From bd3a3635ee67c9af672d2aa2611823d946b86523 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Fri, 10 Apr 2026 10:15:55 -0700 Subject: [PATCH 01/67] refac --- backend/open_webui/utils/auth.py | 47 +++++++++++++++++------------ backend/open_webui/utils/logger.py | 4 ++- src/lib/components/chat/Chat.svelte | 8 +++-- 3 files changed, 36 insertions(+), 23 deletions(-) diff --git a/backend/open_webui/utils/auth.py b/backend/open_webui/utils/auth.py index 16bc36500b..fcdadf9acf 100644 --- a/backend/open_webui/utils/auth.py +++ b/backend/open_webui/utils/auth.py @@ -19,7 +19,6 @@ import pytz from pytz import UTC from typing import Optional, Union, List, Dict -from opentelemetry import trace from open_webui.utils.access_control import has_permission @@ -30,6 +29,7 @@ from open_webui.models.auths import Auths from open_webui.constants import ERROR_MESSAGES from open_webui.env import ( + ENABLE_OTEL, ENABLE_PASSWORD_VALIDATION, OFFLINE_MODE, LICENSE_BLOB, @@ -327,12 +327,15 @@ async def get_current_user( user = get_current_user_by_api_key(request, token) # Add user info to current span - current_span = trace.get_current_span() - if current_span: - current_span.set_attribute('client.user.id', user.id) - current_span.set_attribute('client.user.email', user.email) - current_span.set_attribute('client.user.role', user.role) - current_span.set_attribute('client.auth.type', 'api_key') + if ENABLE_OTEL: + from opentelemetry import trace + + current_span = trace.get_current_span() + if current_span: + current_span.set_attribute('client.user.id', user.id) + current_span.set_attribute('client.user.email', user.email) + current_span.set_attribute('client.user.role', user.role) + current_span.set_attribute('client.auth.type', 'api_key') return user @@ -369,12 +372,15 @@ async def get_current_user( ) # Add user info to current span - current_span = trace.get_current_span() - if current_span: - current_span.set_attribute('client.user.id', user.id) - current_span.set_attribute('client.user.email', user.email) - current_span.set_attribute('client.user.role', user.role) - current_span.set_attribute('client.auth.type', 'jwt') + if ENABLE_OTEL: + from opentelemetry import trace + + current_span = trace.get_current_span() + if current_span: + current_span.set_attribute('client.user.id', user.id) + current_span.set_attribute('client.user.email', user.email) + current_span.set_attribute('client.user.role', user.role) + current_span.set_attribute('client.auth.type', 'jwt') # Refresh the user's last active timestamp asynchronously # to prevent blocking the request @@ -422,12 +428,15 @@ def get_current_user_by_api_key(request, api_key: str): raise HTTPException(status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.API_KEY_NOT_ALLOWED) # Add user info to current span - current_span = trace.get_current_span() - if current_span: - current_span.set_attribute('client.user.id', user.id) - current_span.set_attribute('client.user.email', user.email) - current_span.set_attribute('client.user.role', user.role) - current_span.set_attribute('client.auth.type', 'api_key') + if ENABLE_OTEL: + from opentelemetry import trace + + current_span = trace.get_current_span() + if current_span: + current_span.set_attribute('client.user.id', user.id) + current_span.set_attribute('client.user.email', user.email) + current_span.set_attribute('client.user.role', user.role) + current_span.set_attribute('client.auth.type', 'api_key') Users.update_last_active_by_id(user.id) return user diff --git a/backend/open_webui/utils/logger.py b/backend/open_webui/utils/logger.py index 49b7973c57..fa4e77f53d 100644 --- a/backend/open_webui/utils/logger.py +++ b/backend/open_webui/utils/logger.py @@ -4,7 +4,7 @@ import sys from typing import TYPE_CHECKING from loguru import logger -from opentelemetry import trace + from open_webui.env import ( ENABLE_AUDIT_STDOUT, ENABLE_AUDIT_LOGS_FILE, @@ -100,6 +100,8 @@ class InterceptHandler(logging.Handler): if not ENABLE_OTEL: return {} + from opentelemetry import trace + extras = {} context = trace.get_current_span().get_span_context() if context.is_valid: diff --git a/src/lib/components/chat/Chat.svelte b/src/lib/components/chat/Chat.svelte index 3f0f8d13e4..f1bd5a66e1 100644 --- a/src/lib/components/chat/Chat.svelte +++ b/src/lib/components/chat/Chat.svelte @@ -1245,10 +1245,12 @@ } } - if (query) { - messageInput?.setText(query); + if (query || eventFiles?.length) { + if (query) { + messageInput?.setText(query); + } await tick(); - submitHandler(query); + submitHandler(query || ''); } } else if ($page.url.searchParams.get('q')) { const q = $page.url.searchParams.get('q') ?? ''; From be38ca8e814ce6bd75e489afb456caa9cefcfa35 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Sat, 11 Apr 2026 14:44:05 -0600 Subject: [PATCH 02/67] refac --- src/routes/+layout.svelte | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/src/routes/+layout.svelte b/src/routes/+layout.svelte index b028236533..8b4f5b3e36 100644 --- a/src/routes/+layout.svelte +++ b/src/routes/+layout.svelte @@ -982,6 +982,14 @@ } catch (error) { console.error('Error refreshing backend config:', error); } + + // Relay auth token to desktop app for API access + if (window.electronAPI?.send) { + window.electronAPI.send({ + type: 'token:update', + token: localStorage.token + }).catch(() => {}); + } } else { // Redirect Invalid Session User to /auth Page localStorage.removeItem('token'); From 09dccdd42fa728e06929c4f3f3733094cb302528 Mon Sep 17 00:00:00 2001 From: Skyzi000 <38061609+Skyzi000@users.noreply.github.com> Date: Sun, 12 Apr 2026 05:55:04 +0900 Subject: [PATCH 03/67] fix: use unique client_id per ComfyUI WebSocket connection (#23592) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Using user.id as client_id causes WebSocket deadlocks when the same user generates images concurrently (e.g., multi-model chat). ComfyUI routes messages by clientId, so shared IDs mean only one connection receives the completion — others hang forever. Generate a unique UUID per request, matching ComfyUI's own examples. --- backend/open_webui/routers/images.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/backend/open_webui/routers/images.py b/backend/open_webui/routers/images.py index dca9a58a7a..0e56da560b 100644 --- a/backend/open_webui/routers/images.py +++ b/backend/open_webui/routers/images.py @@ -681,7 +681,7 @@ async def image_generations( res = await comfyui_create_image( model, form_data, - user.id, + str(uuid.uuid4()), request.app.state.config.COMFYUI_BASE_URL, request.app.state.config.COMFYUI_API_KEY, ) @@ -1011,7 +1011,7 @@ async def image_edits( res = await comfyui_edit_image( model, form_data, - user.id, + str(uuid.uuid4()), request.app.state.config.IMAGES_EDIT_COMFYUI_BASE_URL, request.app.state.config.IMAGES_EDIT_COMFYUI_API_KEY, ) From a600f67d6b420136a199ab1efa8793c8e05f912c Mon Sep 17 00:00:00 2001 From: Colin Chen <1207878+silenceroom@users.noreply.github.com> Date: Sun, 12 Apr 2026 04:55:18 +0800 Subject: [PATCH 04/67] i18n: fix Chinese translation for Web Upload permission (#23596) Differentiate between "Allow File Upload" and "Allow Web Upload" in Chinese translations to help administrators understand the distinction: - "Allow File Upload" = local file, cloud storage uploads - "Allow Web Upload" = URL, YouTube, web content uploads --- src/lib/i18n/locales/zh-CN/translation.json | 2 +- src/lib/i18n/locales/zh-TW/translation.json | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/src/lib/i18n/locales/zh-CN/translation.json b/src/lib/i18n/locales/zh-CN/translation.json index 52295330c8..5146725806 100644 --- a/src/lib/i18n/locales/zh-CN/translation.json +++ b/src/lib/i18n/locales/zh-CN/translation.json @@ -137,7 +137,7 @@ "Allow Text to Speech": "允许文本转语音", "Allow User Location": "获取您的位置", "Allow Voice Interruption in Call": "允许语音通话时打断对话", - "Allow Web Upload": "允许上传文件", + "Allow Web Upload": "允许从网络上传内容", "Allowed Endpoints": "允许的接口", "Allowed File Extensions": "允许的文件扩展名", "Allowed file extensions for upload. Separate multiple extensions with commas. Leave empty for all file types.": "文件上传允许的扩展名。多个扩展名用逗号分隔。留空以允许所有文件类型。", diff --git a/src/lib/i18n/locales/zh-TW/translation.json b/src/lib/i18n/locales/zh-TW/translation.json index 497f7a623b..d15e7b486f 100644 --- a/src/lib/i18n/locales/zh-TW/translation.json +++ b/src/lib/i18n/locales/zh-TW/translation.json @@ -137,7 +137,7 @@ "Allow Text to Speech": "允許文字轉語音", "Allow User Location": "允許使用者位置", "Allow Voice Interruption in Call": "允許在通話中打斷語音", - "Allow Web Upload": "允許上傳檔案", + "Allow Web Upload": "允許從網路上傳內容", "Allowed Endpoints": "允許的端點", "Allowed File Extensions": "允許的檔案副檔名", "Allowed file extensions for upload. Separate multiple extensions with commas. Leave empty for all file types.": "允許上傳的檔案副檔名。多個副檔名請用逗號分隔,留空則允許所有檔案類型。", From 69b3ec35116eb279ae9cab35ac30a0546161f974 Mon Sep 17 00:00:00 2001 From: Classic298 <27028174+Classic298@users.noreply.github.com> Date: Sat, 11 Apr 2026 22:57:47 +0200 Subject: [PATCH 05/67] scim: use hmac.compare_digest for bearer token check (#23577) --- backend/open_webui/routers/scim.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/backend/open_webui/routers/scim.py b/backend/open_webui/routers/scim.py index 56923bc447..7bc0157b19 100644 --- a/backend/open_webui/routers/scim.py +++ b/backend/open_webui/routers/scim.py @@ -5,6 +5,7 @@ Provides System for Cross-domain Identity Management endpoints for users and gro NOTE: This is an experimental implementation and may not fully comply with SCIM 2.0 standards, and is subject to change. """ +import hmac import logging import uuid import time @@ -278,7 +279,7 @@ def get_scim_auth(request: Request, authorization: Optional[str] = Header(None)) if hasattr(scim_token, 'value'): scim_token = scim_token.value log.debug(f'SCIM token configured: {bool(scim_token)}') - if not scim_token or token != scim_token: + if not scim_token or not hmac.compare_digest(token, scim_token): raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail='Invalid SCIM token', From bf49358185ca9291a28653654d3b3da7a2481b83 Mon Sep 17 00:00:00 2001 From: Algorithm5838 <108630393+Algorithm5838@users.noreply.github.com> Date: Sun, 12 Apr 2026 00:00:33 +0300 Subject: [PATCH 06/67] refactor: use shared unescapeHtml in CodeBlock (#23553) --- .../components/chat/Messages/CodeBlock.svelte | 18 +++--------------- 1 file changed, 3 insertions(+), 15 deletions(-) diff --git a/src/lib/components/chat/Messages/CodeBlock.svelte b/src/lib/components/chat/Messages/CodeBlock.svelte index 308b15aa9c..8bcd241e8e 100644 --- a/src/lib/components/chat/Messages/CodeBlock.svelte +++ b/src/lib/components/chat/Messages/CodeBlock.svelte @@ -10,7 +10,8 @@ copyToClipboard, initMermaid, renderMermaidDiagram, - renderVegaVisualization + renderVegaVisualization, + unescapeHtml } from '$lib/utils'; import 'highlight.js/styles/github-dark.min.css'; @@ -405,21 +406,8 @@ const onAttributesUpdate = () => { if (attributes?.output) { - // Create a helper function to unescape HTML entities - const unescapeHtml = (html) => { - const textArea = document.createElement('textarea'); - textArea.innerHTML = html; - return textArea.value; - }; - try { - // Unescape the HTML-encoded string - const unescapedOutput = unescapeHtml(attributes.output); - - // Parse the unescaped string into JSON - const output = JSON.parse(unescapedOutput); - - // Assign the parsed values to variables + const output = JSON.parse(unescapeHtml(attributes.output)); stdout = output.stdout; stderr = output.stderr; result = output.result; From b6db719758fe3e3cd167dcd2e290c5582a462b5f Mon Sep 17 00:00:00 2001 From: Algorithm5838 <108630393+Algorithm5838@users.noreply.github.com> Date: Sun, 12 Apr 2026 00:02:17 +0300 Subject: [PATCH 07/67] perf: build mention regex once in factory closure (#23551) --- src/lib/utils/marked/mention-extension.ts | 37 +++++++++++------------ 1 file changed, 17 insertions(+), 20 deletions(-) diff --git a/src/lib/utils/marked/mention-extension.ts b/src/lib/utils/marked/mention-extension.ts index 4e3865fd89..b56edfbee3 100644 --- a/src/lib/utils/marked/mention-extension.ts +++ b/src/lib/utils/marked/mention-extension.ts @@ -18,24 +18,6 @@ function mentionStart(src: string) { return src.indexOf('<'); } -function mentionTokenizer(this: any, src: string, options: MentionOptions = {}) { - const trigger = options.triggerChar ?? '@'; - // Build dynamic regex for `<@id>`, `<@id|label>`, `<@id|>` - // Added forward slash (/) to the character class for IDs - const re = new RegExp(`^<\\${trigger}([\\w.\\-:/]+)(?:\\|([^>]*))?>`); - const m = re.exec(src); - if (!m) return; - - const [, id, label] = m; - return { - type: 'mention', - raw: m[0], - triggerChar: trigger, - id, - label: label && label.length > 0 ? label : id - }; -} - function mentionRenderer(token: any, options: MentionOptions = {}) { const trigger = options.triggerChar ?? '@'; const cls = options.className ?? 'mention'; @@ -55,15 +37,30 @@ function mentionRenderer(token: any, options: MentionOptions = {}) { } export function mentionExtension(opts: MentionOptions = {}) { + // Compile the regex once when the extension is created, not on every tokenizer call. + // mentionStart fires on every '<' in the document, making the tokenizer a hot path. + const trigger = opts.triggerChar ?? '@'; + const re = new RegExp(`^<\\${trigger}([\\w.\\-:/]+)(?:\\|([^>]*))?>`); + const snapshot: MentionOptions = { triggerChar: trigger, className: opts.className, extraAttrs: opts.extraAttrs }; + return { name: 'mention', level: 'inline' as const, start: mentionStart, tokenizer(src: string) { - return mentionTokenizer.call(this, src, opts); + const m = re.exec(src); + if (!m) return; + const [, id, label] = m; + return { + type: 'mention', + raw: m[0], + triggerChar: trigger, + id, + label: label && label.length > 0 ? label : id + }; }, renderer(token: any) { - return mentionRenderer(token, opts); + return mentionRenderer(token, snapshot); } }; } From 6acaaea59a50ec26da03e6144017a2fd86241ce9 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Sat, 11 Apr 2026 15:23:37 -0600 Subject: [PATCH 08/67] refac --- backend/open_webui/routers/channels.py | 6 ++++-- backend/open_webui/routers/files.py | 2 +- 2 files changed, 5 insertions(+), 3 deletions(-) diff --git a/backend/open_webui/routers/channels.py b/backend/open_webui/routers/channels.py index aa1ea52662..5241fe7445 100644 --- a/backend/open_webui/routers/channels.py +++ b/backend/open_webui/routers/channels.py @@ -1162,10 +1162,11 @@ async def get_channel_message( if message.channel_id != id: raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.DEFAULT()) + message_user = Users.get_user_by_id(message.user_id, db=db) return MessageResponse( **{ **message.model_dump(), - 'user': UserNameResponse(**Users.get_user_by_id(message.user_id, db=db).model_dump()), + 'user': UserNameResponse(**message_user.model_dump()) if message_user else None, } ) @@ -1245,10 +1246,11 @@ async def pin_channel_message( try: Messages.update_is_pinned_by_id(message_id, form_data.is_pinned, user.id, db=db) message = Messages.get_message_by_id(message_id, db=db) + message_user = Users.get_user_by_id(message.user_id, db=db) return MessageUserResponse( **{ **message.model_dump(), - 'user': UserNameResponse(**Users.get_user_by_id(message.user_id, db=db).model_dump()), + 'user': UserNameResponse(**message_user.model_dump()) if message_user else None, } ) except Exception as e: diff --git a/backend/open_webui/routers/files.py b/backend/open_webui/routers/files.py index 48172e744e..84227c0eca 100644 --- a/backend/open_webui/routers/files.py +++ b/backend/open_webui/routers/files.py @@ -656,7 +656,7 @@ async def get_html_file_content_by_id(id: str, user=Depends(get_verified_user), ) file_user = Users.get_user_by_id(file.user_id, db=db) - if not file_user.role == 'admin': + if not file_user or file_user.role != 'admin': raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND, From faf935ef5285b8eaa816b556354e145a4a70c7ee Mon Sep 17 00:00:00 2001 From: Classic298 <27028174+Classic298@users.noreply.github.com> Date: Sat, 11 Apr 2026 23:31:34 +0200 Subject: [PATCH 09/67] auths: match JWT expiry on /auths/add with other sign-in paths (#23576) --- backend/open_webui/routers/auths.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/backend/open_webui/routers/auths.py b/backend/open_webui/routers/auths.py index 88f0fe69fb..8cca63baf0 100644 --- a/backend/open_webui/routers/auths.py +++ b/backend/open_webui/routers/auths.py @@ -877,7 +877,8 @@ async def add_user( db=db, ) - token = create_token(data={'id': user.id}) + expires_delta = parse_duration(request.app.state.config.JWT_EXPIRES_IN) + token = create_token(data={'id': user.id}, expires_delta=expires_delta) return { 'token': token, 'token_type': 'Bearer', From c0ac10d5db5f360feb9be456c6367cca59fcfc4d Mon Sep 17 00:00:00 2001 From: Classic298 <27028174+Classic298@users.noreply.github.com> Date: Sat, 11 Apr 2026 23:32:05 +0200 Subject: [PATCH 10/67] fix: honor REDIS_SOCKET_CONNECT_TIMEOUT on non-sentinel clients (#23572) * fix(redis): honor REDIS_SOCKET_CONNECT_TIMEOUT on non-sentinel clients Previously only the sentinel path passed REDIS_SOCKET_CONNECT_TIMEOUT through to the Redis client. Plain redis:// and cluster URLs fell back to redis-py's default (no explicit connect timeout), so a hung Redis or a black-holed network path could stall the whole worker until the kernel gave up. Forwarding the same env var to from_url()/RedisCluster keeps the behavior consistent across all deployment topologies. * fix(redis): gate socket_connect_timeout on is-not-None, not truthiness Addresses review feedback: the truthiness check on REDIS_SOCKET_CONNECT_TIMEOUT silently dropped an explicit 0 value and was inconsistent with the sentinel construction path, which forwards the value directly. Switch to `is not None` so any user-configured value (including 0) is passed through to from_url() and RedisCluster.from_url(). --------- Co-authored-by: Claude --- backend/open_webui/utils/redis.py | 30 ++++++++++++++++++++++++++---- 1 file changed, 26 insertions(+), 4 deletions(-) diff --git a/backend/open_webui/utils/redis.py b/backend/open_webui/utils/redis.py index 55d08147a9..c2e5da1fae 100644 --- a/backend/open_webui/utils/redis.py +++ b/backend/open_webui/utils/redis.py @@ -191,6 +191,12 @@ def get_redis_connection( connection = None + connect_timeout_kwargs = ( + {'socket_connect_timeout': REDIS_SOCKET_CONNECT_TIMEOUT} + if REDIS_SOCKET_CONNECT_TIMEOUT is not None + else {} + ) + if async_mode: import redis.asyncio as redis @@ -214,9 +220,17 @@ def get_redis_connection( elif redis_cluster: if not redis_url: raise ValueError('Redis URL must be provided for cluster mode.') - return redis.cluster.RedisCluster.from_url(redis_url, decode_responses=decode_responses) + return redis.cluster.RedisCluster.from_url( + redis_url, + decode_responses=decode_responses, + **connect_timeout_kwargs, + ) elif redis_url: - connection = redis.from_url(redis_url, decode_responses=decode_responses) + connection = redis.from_url( + redis_url, + decode_responses=decode_responses, + **connect_timeout_kwargs, + ) else: import redis @@ -239,9 +253,17 @@ def get_redis_connection( elif redis_cluster: if not redis_url: raise ValueError('Redis URL must be provided for cluster mode.') - return redis.cluster.RedisCluster.from_url(redis_url, decode_responses=decode_responses) + return redis.cluster.RedisCluster.from_url( + redis_url, + decode_responses=decode_responses, + **connect_timeout_kwargs, + ) elif redis_url: - connection = redis.Redis.from_url(redis_url, decode_responses=decode_responses) + connection = redis.Redis.from_url( + redis_url, + decode_responses=decode_responses, + **connect_timeout_kwargs, + ) _CONNECTION_CACHE[cache_key] = connection return connection From 008cd8e6b9699b7fb68d520ef3a399b7f239fd5a Mon Sep 17 00:00:00 2001 From: Wang Weixuan Date: Sun, 12 Apr 2026 05:36:22 +0800 Subject: [PATCH 11/67] refac: import fastapi instrumentor directly (#23530) --- backend/open_webui/utils/telemetry/instrumentors.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/backend/open_webui/utils/telemetry/instrumentors.py b/backend/open_webui/utils/telemetry/instrumentors.py index 394e7178d6..fe8e9ba799 100644 --- a/backend/open_webui/utils/telemetry/instrumentors.py +++ b/backend/open_webui/utils/telemetry/instrumentors.py @@ -7,8 +7,8 @@ from aiohttp import ( TraceRequestEndParams, TraceRequestExceptionParams, ) -from chromadb.telemetry.opentelemetry.fastapi import instrument_fastapi from fastapi import FastAPI +from opentelemetry.instrumentation.fastapi import FastAPIInstrumentor from opentelemetry.instrumentation.httpx import ( HTTPXClientInstrumentor, RequestInfo, @@ -176,7 +176,7 @@ class Instrumentor(BaseInstrumentor): return [] def _instrument(self, **kwargs): - instrument_fastapi(app=self.app) + FastAPIInstrumentor.instrument_app(app=self.app) SQLAlchemyInstrumentor().instrument(engine=self.db_engine) RedisInstrumentor().instrument(request_hook=redis_request_hook) RequestsInstrumentor().instrument(request_hook=requests_hook, response_hook=response_hook) From aacf95cf76ca67deefe8b7a7baf6e45c5e737d38 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Sat, 11 Apr 2026 16:08:16 -0600 Subject: [PATCH 12/67] refac --- src/lib/components/chat/Chat.svelte | 48 +++++++++++++++++------------ src/lib/stores/index.ts | 4 +-- src/routes/+layout.svelte | 7 ++++- 3 files changed, 36 insertions(+), 23 deletions(-) diff --git a/src/lib/components/chat/Chat.svelte b/src/lib/components/chat/Chat.svelte index f1bd5a66e1..5b96f82c2f 100644 --- a/src/lib/components/chat/Chat.svelte +++ b/src/lib/components/chat/Chat.svelte @@ -1226,31 +1226,39 @@ showControls.set(true); } - // Consume one-shot desktop event (e.g. Spotlight query + attachments) + // Consume one-shot desktop event (e.g. Spotlight query, call shortcut) if ($desktopEvent) { - const { query, files: eventFiles } = $desktopEvent; + const event = $desktopEvent; desktopEvent.set(null); - // Attach screenshot images from desktop (e.g. Spotlight region capture) - if (eventFiles?.length) { - for (const ef of eventFiles) { - files = [ - ...files, - { - type: 'image', - url: ef.dataUrl, - name: ef.name - } - ]; - } - } + if (event.type === 'call') { + showCallOverlay.set(true); + showControls.set(true); + } else if (event.type === 'query') { + const query = event.data?.query; + const eventFiles = event.data?.files; - if (query || eventFiles?.length) { - if (query) { - messageInput?.setText(query); + // Attach screenshot images from desktop (e.g. Spotlight region capture) + if (eventFiles?.length) { + for (const ef of eventFiles) { + files = [ + ...files, + { + type: 'image', + url: ef.dataUrl, + name: ef.name + } + ]; + } + } + + if (query || eventFiles?.length) { + if (query) { + messageInput?.setText(query); + } + await tick(); + submitHandler(query || ''); } - await tick(); - submitHandler(query || ''); } } else if ($page.url.searchParams.get('q')) { const q = $page.url.searchParams.get('q') ?? ''; diff --git a/src/lib/stores/index.ts b/src/lib/stores/index.ts index edcea62137..05bd83ea1f 100644 --- a/src/lib/stores/index.ts +++ b/src/lib/stores/index.ts @@ -116,8 +116,8 @@ export const temporaryChatEnabled = writable(false); // Set by +layout.svelte, consumed and cleared by Chat.svelte. export type DesktopEventFile = { name: string; mimeType: string; dataUrl: string }; export type DesktopEvent = { - query?: string; - files?: DesktopEventFile[]; + type: string; + data?: any; }; export const desktopEvent: Writable = writable(null); export const scrollPaginationEnabled = writable(false); diff --git a/src/routes/+layout.svelte b/src/routes/+layout.svelte index 8b4f5b3e36..252e1cf7a8 100644 --- a/src/routes/+layout.svelte +++ b/src/routes/+layout.svelte @@ -712,7 +712,12 @@ return; } if (event.type === 'query' && (event.data?.query || event.data?.files?.length)) { - desktopEvent.set({ query: event.data.query, files: event.data.files }); + desktopEvent.set(event); + await goto('/'); + return; + } + if (event.type === 'call') { + desktopEvent.set(event); await goto('/'); return; } From db7f122cb0dc2ae519e5c77e73ad09af725a2363 Mon Sep 17 00:00:00 2001 From: Classic298 <27028174+Classic298@users.noreply.github.com> Date: Sun, 12 Apr 2026 00:09:12 +0200 Subject: [PATCH 13/67] fix(redis): add opt-in TCP socket keepalive on all client connections (#23571) Introduces REDIS_SOCKET_KEEPALIVE and wires socket_keepalive=True through to every Redis client created by get_redis_connection (plain, cluster and sentinel paths, sync and async). When enabled, the kernel sends TCP keepalive probes on idle connections so half-closed sockets (e.g. after a silent firewall/LB reset or a NIC flap) are detected before the next command lands on them and the request never sees a "Connection reset by peer" error. Defaults to off so existing deployments see no behavioural change. Operators who want the protection set REDIS_SOCKET_KEEPALIVE=true in their environment. Co-authored-by: Claude --- backend/open_webui/env.py | 9 +++++++++ backend/open_webui/utils/redis.py | 11 +++++++++++ 2 files changed, 20 insertions(+) diff --git a/backend/open_webui/env.py b/backend/open_webui/env.py index 53737a1f2f..f612248ffe 100644 --- a/backend/open_webui/env.py +++ b/backend/open_webui/env.py @@ -426,6 +426,15 @@ try: except ValueError: REDIS_SOCKET_CONNECT_TIMEOUT = None +# Whether to enable TCP SO_KEEPALIVE on Redis client sockets. Opt-in: +# defaults to off so behavior is unchanged for existing deployments. When +# enabled, the kernel sends TCP keepalive probes on idle connections so +# half-closed sockets (e.g. after a silent firewall/LB reset or a NIC +# flap) are detected before the next command lands on them. +REDIS_SOCKET_KEEPALIVE = ( + os.environ.get('REDIS_SOCKET_KEEPALIVE', 'False').lower() == 'true' +) + REDIS_RECONNECT_DELAY = os.environ.get('REDIS_RECONNECT_DELAY', '') if REDIS_RECONNECT_DELAY == '': diff --git a/backend/open_webui/utils/redis.py b/backend/open_webui/utils/redis.py index c2e5da1fae..ec1dee5e9b 100644 --- a/backend/open_webui/utils/redis.py +++ b/backend/open_webui/utils/redis.py @@ -10,6 +10,7 @@ import redis from open_webui.env import ( REDIS_CLUSTER, REDIS_SOCKET_CONNECT_TIMEOUT, + REDIS_SOCKET_KEEPALIVE, REDIS_SENTINEL_HOSTS, REDIS_SENTINEL_MAX_RETRY_COUNT, REDIS_SENTINEL_PORT, @@ -197,6 +198,10 @@ def get_redis_connection( else {} ) + keepalive_kwargs = ( + {'socket_keepalive': True} if REDIS_SOCKET_KEEPALIVE else {} + ) + if async_mode: import redis.asyncio as redis @@ -211,6 +216,7 @@ def get_redis_connection( password=redis_config['password'], decode_responses=decode_responses, socket_connect_timeout=REDIS_SOCKET_CONNECT_TIMEOUT, + **keepalive_kwargs, ) connection = SentinelRedisProxy( sentinel, @@ -224,12 +230,14 @@ def get_redis_connection( redis_url, decode_responses=decode_responses, **connect_timeout_kwargs, + **keepalive_kwargs, ) elif redis_url: connection = redis.from_url( redis_url, decode_responses=decode_responses, **connect_timeout_kwargs, + **keepalive_kwargs, ) else: import redis @@ -244,6 +252,7 @@ def get_redis_connection( password=redis_config['password'], decode_responses=decode_responses, socket_connect_timeout=REDIS_SOCKET_CONNECT_TIMEOUT, + **keepalive_kwargs, ) connection = SentinelRedisProxy( sentinel, @@ -257,12 +266,14 @@ def get_redis_connection( redis_url, decode_responses=decode_responses, **connect_timeout_kwargs, + **keepalive_kwargs, ) elif redis_url: connection = redis.Redis.from_url( redis_url, decode_responses=decode_responses, **connect_timeout_kwargs, + **keepalive_kwargs, ) _CONNECTION_CACHE[cache_key] = connection From 588b81eedaacbfd7394b707ae1600d9fb729b809 Mon Sep 17 00:00:00 2001 From: Classic298 <27028174+Classic298@users.noreply.github.com> Date: Sun, 12 Apr 2026 00:17:19 +0200 Subject: [PATCH 14/67] fix(redis): add opt-in health_check_interval for stale pooled connections (#23573) Introduces REDIS_HEALTH_CHECK_INTERVAL and wires it through to every Redis client created by get_redis_connection (plain, cluster and sentinel paths, sync and async). When set, redis-py will PING any connection idle longer than the interval on checkout, so dead sockets are surfaced as reconnectable errors before a real command lands on them. Defaults to unset (empty string) so existing deployments see no behavioural change. Operators who want the protection should set it shorter than the Redis server `timeout` setting and any firewall/LB idle timeout on the path to Redis. Co-authored-by: Claude --- backend/open_webui/env.py | 14 ++++++++++++++ backend/open_webui/utils/redis.py | 13 +++++++++++++ 2 files changed, 27 insertions(+) diff --git a/backend/open_webui/env.py b/backend/open_webui/env.py index f612248ffe..1695cec20a 100644 --- a/backend/open_webui/env.py +++ b/backend/open_webui/env.py @@ -435,6 +435,20 @@ REDIS_SOCKET_KEEPALIVE = ( os.environ.get('REDIS_SOCKET_KEEPALIVE', 'False').lower() == 'true' ) +# How often (in seconds) redis-py should PING an idle pooled connection +# before reusing it. Opt-in: defaults to unset (empty string) so behavior +# is unchanged for existing deployments. When set, should be shorter than +# the Redis server `timeout` setting and any firewall/LB idle timeout on +# the path to Redis, so stale connections are detected before a real +# command lands on them. Set to 0 or empty to disable. +REDIS_HEALTH_CHECK_INTERVAL = os.environ.get('REDIS_HEALTH_CHECK_INTERVAL', '') +try: + REDIS_HEALTH_CHECK_INTERVAL = int(REDIS_HEALTH_CHECK_INTERVAL) + if REDIS_HEALTH_CHECK_INTERVAL <= 0: + REDIS_HEALTH_CHECK_INTERVAL = None +except ValueError: + REDIS_HEALTH_CHECK_INTERVAL = None + REDIS_RECONNECT_DELAY = os.environ.get('REDIS_RECONNECT_DELAY', '') if REDIS_RECONNECT_DELAY == '': diff --git a/backend/open_webui/utils/redis.py b/backend/open_webui/utils/redis.py index ec1dee5e9b..61a4d74c7c 100644 --- a/backend/open_webui/utils/redis.py +++ b/backend/open_webui/utils/redis.py @@ -9,6 +9,7 @@ import redis from open_webui.env import ( REDIS_CLUSTER, + REDIS_HEALTH_CHECK_INTERVAL, REDIS_SOCKET_CONNECT_TIMEOUT, REDIS_SOCKET_KEEPALIVE, REDIS_SENTINEL_HOSTS, @@ -202,6 +203,12 @@ def get_redis_connection( {'socket_keepalive': True} if REDIS_SOCKET_KEEPALIVE else {} ) + health_check_kwargs = ( + {'health_check_interval': REDIS_HEALTH_CHECK_INTERVAL} + if REDIS_HEALTH_CHECK_INTERVAL + else {} + ) + if async_mode: import redis.asyncio as redis @@ -217,6 +224,7 @@ def get_redis_connection( decode_responses=decode_responses, socket_connect_timeout=REDIS_SOCKET_CONNECT_TIMEOUT, **keepalive_kwargs, + **health_check_kwargs, ) connection = SentinelRedisProxy( sentinel, @@ -231,6 +239,7 @@ def get_redis_connection( decode_responses=decode_responses, **connect_timeout_kwargs, **keepalive_kwargs, + **health_check_kwargs, ) elif redis_url: connection = redis.from_url( @@ -238,6 +247,7 @@ def get_redis_connection( decode_responses=decode_responses, **connect_timeout_kwargs, **keepalive_kwargs, + **health_check_kwargs, ) else: import redis @@ -253,6 +263,7 @@ def get_redis_connection( decode_responses=decode_responses, socket_connect_timeout=REDIS_SOCKET_CONNECT_TIMEOUT, **keepalive_kwargs, + **health_check_kwargs, ) connection = SentinelRedisProxy( sentinel, @@ -267,6 +278,7 @@ def get_redis_connection( decode_responses=decode_responses, **connect_timeout_kwargs, **keepalive_kwargs, + **health_check_kwargs, ) elif redis_url: connection = redis.Redis.from_url( @@ -274,6 +286,7 @@ def get_redis_connection( decode_responses=decode_responses, **connect_timeout_kwargs, **keepalive_kwargs, + **health_check_kwargs, ) _CONNECTION_CACHE[cache_key] = connection From 674695918e5e3e1811314ce2a082c5bbb42d76b2 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Sat, 11 Apr 2026 16:44:12 -0600 Subject: [PATCH 15/67] refac --- backend/open_webui/tools/builtin.py | 329 ++++++++++++++++++ backend/open_webui/utils/tools.py | 48 ++- .../workspace/Models/BuiltinTools.svelte | 4 + 3 files changed, 376 insertions(+), 5 deletions(-) diff --git a/backend/open_webui/tools/builtin.py b/backend/open_webui/tools/builtin.py index 1823e71df2..d6032eb900 100644 --- a/backend/open_webui/tools/builtin.py +++ b/backend/open_webui/tools/builtin.py @@ -2494,3 +2494,332 @@ async def tasks( except Exception as e: log.exception(f'tasks error: {e}') return json.dumps({'error': str(e)}) + + +# ============================================================================= +# AUTOMATION TOOLS +# ============================================================================= + + +async def create_automation( + name: str, + prompt: str, + rrule: str, + model_id: Optional[str] = None, + __request__: Request = None, + __user__: dict = None, + __metadata__: dict = None, +) -> str: + """ + Create a scheduled automation that runs a prompt on a recurring or one-time schedule. + Use this when the user wants to schedule a task to run automatically. + + The rrule parameter must be a valid iCalendar RRULE string. Common examples: + - Every day at 9am: "DTSTART:20250101T090000\\nRRULE:FREQ=DAILY" + - Every Monday at 8am: "DTSTART:20250106T080000\\nRRULE:FREQ=WEEKLY;BYDAY=MO" + - Every hour: "RRULE:FREQ=HOURLY;INTERVAL=1" + - Every 30 minutes: "RRULE:FREQ=MINUTELY;INTERVAL=30" + - Once at a specific time: "DTSTART:20250415T140000\\nRRULE:FREQ=DAILY;COUNT=1" + - First day of every month: "DTSTART:20250101T090000\\nRRULE:FREQ=MONTHLY;BYMONTHDAY=1" + + The DTSTART time should reflect the desired execution time. Use COUNT=1 for one-time automations. + + :param name: A short descriptive name for the automation + :param prompt: The prompt/instructions to execute on each run + :param rrule: An iCalendar RRULE string defining the schedule + :param model_id: Optional model ID to use. Defaults to the current chat model if omitted. + :return: JSON with the created automation details including id, next scheduled runs + """ + if __request__ is None: + return json.dumps({'error': 'Request context not available'}) + + if not __user__: + return json.dumps({'error': 'User context not available'}) + + try: + from open_webui.models.automations import Automations, AutomationForm, AutomationData + from open_webui.models.users import Users + from open_webui.utils.automations import validate_rrule, next_run_ns, next_n_runs_ns + + user_id = __user__.get('id') + user = Users.get_user_by_id(user_id) + if not user: + return json.dumps({'error': 'User not found'}) + + # Default to current chat's model if not specified + if not model_id: + model_id = (__metadata__ or {}).get('model_id') or (__metadata__ or {}).get('model') + if not model_id: + return json.dumps({'error': 'model_id is required (could not detect current model)'}) + + # Validate the RRULE + try: + validate_rrule(rrule) + except ValueError as e: + return json.dumps({'error': f'Invalid schedule: {e}'}) + + tz = user.timezone + form = AutomationForm( + name=name, + data=AutomationData( + prompt=prompt, + model_id=model_id, + rrule=rrule, + ), + is_active=True, + ) + + automation = Automations.insert(user_id, form, next_run_ns(rrule, tz=tz)) + + return json.dumps( + { + 'status': 'success', + 'id': automation.id, + 'name': automation.name, + 'model_id': model_id, + 'is_active': automation.is_active, + 'next_runs': next_n_runs_ns(rrule, tz=tz), + }, + ensure_ascii=False, + ) + except Exception as e: + log.exception(f'create_automation error: {e}') + return json.dumps({'error': str(e)}) + + +async def update_automation( + automation_id: str, + name: Optional[str] = None, + prompt: Optional[str] = None, + rrule: Optional[str] = None, + model_id: Optional[str] = None, + __request__: Request = None, + __user__: dict = None, +) -> str: + """ + Update an existing automation. Only the provided fields are changed; omitted fields stay the same. + + :param automation_id: The ID of the automation to update + :param name: New name for the automation (optional) + :param prompt: New prompt/instructions (optional) + :param rrule: New iCalendar RRULE schedule string (optional). See create_automation for format examples. + :param model_id: New model ID to use (optional) + :return: JSON with the updated automation details + """ + if __request__ is None: + return json.dumps({'error': 'Request context not available'}) + + if not __user__: + return json.dumps({'error': 'User context not available'}) + + try: + from open_webui.models.automations import Automations, AutomationForm, AutomationData + from open_webui.models.users import Users + from open_webui.utils.automations import validate_rrule, next_run_ns, next_n_runs_ns + + user_id = __user__.get('id') + user = Users.get_user_by_id(user_id) + + automation = Automations.get_by_id(automation_id) + if not automation: + return json.dumps({'error': 'Automation not found'}) + if automation.user_id != user_id: + return json.dumps({'error': 'Access denied'}) + + # Merge provided fields with existing values + new_name = name if name is not None else automation.name + new_prompt = prompt if prompt is not None else automation.data.get('prompt', '') + new_model_id = model_id if model_id is not None else automation.data.get('model_id', '') + new_rrule = rrule if rrule is not None else automation.data.get('rrule', '') + + # Validate RRULE if changed + if rrule is not None: + try: + validate_rrule(new_rrule) + except ValueError as e: + return json.dumps({'error': f'Invalid schedule: {e}'}) + + tz = user.timezone if user else None + form = AutomationForm( + name=new_name, + data=AutomationData( + prompt=new_prompt, + model_id=new_model_id, + rrule=new_rrule, + ), + is_active=automation.is_active, + ) + + updated = Automations.update(automation_id, form, next_run_ns(new_rrule, tz=tz)) + + return json.dumps( + { + 'status': 'success', + 'id': updated.id, + 'name': updated.name, + 'model_id': new_model_id, + 'is_active': updated.is_active, + 'next_runs': next_n_runs_ns(new_rrule, tz=tz), + }, + ensure_ascii=False, + ) + except Exception as e: + log.exception(f'update_automation error: {e}') + return json.dumps({'error': str(e)}) + + +async def list_automations( + status: Optional[str] = None, + count: int = 10, + __request__: Request = None, + __user__: dict = None, +) -> str: + """ + List the user's scheduled automations. + + :param status: Filter by status: "active", "paused", or omit for all + :param count: Maximum number of automations to return (default: 10) + :return: JSON list of automations with id, name, prompt snippet, schedule, status, and next runs + """ + if __request__ is None: + return json.dumps({'error': 'Request context not available'}) + + if not __user__: + return json.dumps({'error': 'User context not available'}) + + try: + from open_webui.models.automations import Automations + from open_webui.models.users import Users + from open_webui.utils.automations import next_n_runs_ns + + user_id = __user__.get('id') + user = Users.get_user_by_id(user_id) + + result = Automations.search_automations( + user_id=user_id, + status=status, + skip=0, + limit=count, + ) + + automations = [] + for item in result.items: + rrule = item.data.get('rrule', '') + prompt_text = item.data.get('prompt', '') + snippet = prompt_text[:100] + ('...' if len(prompt_text) > 100 else '') + + automations.append( + { + 'id': item.id, + 'name': item.name, + 'prompt_snippet': snippet, + 'model_id': item.data.get('model_id', ''), + 'rrule': rrule, + 'is_active': item.is_active, + 'last_run_at': item.last_run_at, + 'next_runs': next_n_runs_ns(rrule, tz=user.timezone if user else None), + } + ) + + return json.dumps( + {'automations': automations, 'total': result.total}, + ensure_ascii=False, + ) + except Exception as e: + log.exception(f'list_automations error: {e}') + return json.dumps({'error': str(e)}) + + +async def toggle_automation( + automation_id: str, + __request__: Request = None, + __user__: dict = None, +) -> str: + """ + Pause or resume a scheduled automation. If active, it will be paused. If paused, it will be resumed. + + :param automation_id: The ID of the automation to toggle + :return: JSON with the updated automation status + """ + if __request__ is None: + return json.dumps({'error': 'Request context not available'}) + + if not __user__: + return json.dumps({'error': 'User context not available'}) + + try: + from open_webui.models.automations import Automations + from open_webui.models.users import Users + from open_webui.utils.automations import next_run_ns + + user_id = __user__.get('id') + user = Users.get_user_by_id(user_id) + + automation = Automations.get_by_id(automation_id) + if not automation: + return json.dumps({'error': 'Automation not found'}) + if automation.user_id != user_id: + return json.dumps({'error': 'Access denied'}) + + rrule = automation.data.get('rrule', '') + toggled = Automations.toggle( + automation_id, + next_run_ns(rrule, tz=user.timezone if user else None), + ) + + return json.dumps( + { + 'status': 'success', + 'id': toggled.id, + 'name': toggled.name, + 'is_active': toggled.is_active, + }, + ensure_ascii=False, + ) + except Exception as e: + log.exception(f'toggle_automation error: {e}') + return json.dumps({'error': str(e)}) + + +async def delete_automation( + automation_id: str, + __request__: Request = None, + __user__: dict = None, +) -> str: + """ + Delete a scheduled automation and all its run history. + + :param automation_id: The ID of the automation to delete + :return: JSON confirming the automation was deleted + """ + if __request__ is None: + return json.dumps({'error': 'Request context not available'}) + + if not __user__: + return json.dumps({'error': 'User context not available'}) + + try: + from open_webui.models.automations import Automations, AutomationRuns + + user_id = __user__.get('id') + + automation = Automations.get_by_id(automation_id) + if not automation: + return json.dumps({'error': 'Automation not found'}) + if automation.user_id != user_id: + return json.dumps({'error': 'Access denied'}) + + name = automation.name + AutomationRuns.delete_by_automation(automation_id) + Automations.delete(automation_id) + + return json.dumps( + { + 'status': 'success', + 'message': f'Automation "{name}" deleted', + }, + ensure_ascii=False, + ) + except Exception as e: + log.exception(f'delete_automation error: {e}') + return json.dumps({'error': str(e)}) diff --git a/backend/open_webui/utils/tools.py b/backend/open_webui/utils/tools.py index 377a81d749..aff33caa4d 100644 --- a/backend/open_webui/utils/tools.py +++ b/backend/open_webui/utils/tools.py @@ -86,9 +86,15 @@ from open_webui.tools.builtin import ( view_knowledge_file, view_skill, tasks, + create_automation, + update_automation, + list_automations, + toggle_automation, + delete_automation, ) import copy +from open_webui.utils.access_control import has_permission log = logging.getLogger(__name__) @@ -397,6 +403,18 @@ def get_builtin_tools( builtin_tools = model.get('info', {}).get('meta', {}).get('builtinTools', {}) return builtin_tools.get(category, True) + # Helper to check user-level feature permission (admins always pass) + user = extra_params.get('__user__', {}) + + def has_user_permission(feature_key: str) -> bool: + if user.get('role') == 'admin': + return True + return has_permission( + user.get('id', ''), + f'features.{feature_key}', + request.app.state.config.USER_PERMISSIONS, + ) + # Time utilities - available for date calculations if is_builtin_tool_enabled('time'): builtin_functions.extend([get_current_timestamp, calculate_timestamp]) @@ -440,7 +458,11 @@ def get_builtin_tools( builtin_functions.extend([search_chats, view_chat]) # Add memory tools if builtin category enabled AND enabled for this chat - if is_builtin_tool_enabled('memory') and (features.get('memory') or get_model_capability('memory', False)): + if ( + is_builtin_tool_enabled('memory') + and (features.get('memory') or get_model_capability('memory', False)) + and has_user_permission('memories') + ): builtin_functions.extend( [ search_memories, @@ -457,6 +479,7 @@ def get_builtin_tools( and getattr(request.app.state.config, 'ENABLE_WEB_SEARCH', False) and get_model_capability('web_search') and features.get('web_search') + and has_user_permission('web_search') ): builtin_functions.extend([search_web, fetch_url]) @@ -466,6 +489,7 @@ def get_builtin_tools( and getattr(request.app.state.config, 'ENABLE_IMAGE_GENERATION', False) and get_model_capability('image_generation') and features.get('image_generation') + and has_user_permission('image_generation') ): builtin_functions.append(generate_image) if ( @@ -473,6 +497,7 @@ def get_builtin_tools( and getattr(request.app.state.config, 'ENABLE_IMAGE_EDIT', False) and get_model_capability('image_generation') and features.get('image_generation') + and has_user_permission('image_generation') ): builtin_functions.append(edit_image) @@ -482,15 +507,24 @@ def get_builtin_tools( and getattr(request.app.state.config, 'ENABLE_CODE_INTERPRETER', True) and get_model_capability('code_interpreter') and features.get('code_interpreter') + and has_user_permission('code_interpreter') ): builtin_functions.append(execute_code) - # Notes tools - search, view, create, and update user's notes (if builtin category enabled AND notes enabled globally) - if is_builtin_tool_enabled('notes') and getattr(request.app.state.config, 'ENABLE_NOTES', False): + # Notes tools - search, view, create, and update user's notes + if ( + is_builtin_tool_enabled('notes') + and getattr(request.app.state.config, 'ENABLE_NOTES', False) + and has_user_permission('notes') + ): builtin_functions.extend([search_notes, view_note, write_note, replace_note_content]) - # Channels tools - search channels and messages (if builtin category enabled AND channels enabled globally) - if is_builtin_tool_enabled('channels') and getattr(request.app.state.config, 'ENABLE_CHANNELS', False): + # Channels tools - search channels and messages + if ( + is_builtin_tool_enabled('channels') + and getattr(request.app.state.config, 'ENABLE_CHANNELS', False) + and has_user_permission('channels') + ): builtin_functions.extend( [ search_channels, @@ -508,6 +542,10 @@ def get_builtin_tools( if is_builtin_tool_enabled('tasks'): builtin_functions.append(tasks) + # Automation tools - create and manage scheduled automations from chat + if is_builtin_tool_enabled('automations') and has_user_permission('automations'): + builtin_functions.extend([create_automation, update_automation, list_automations, toggle_automation, delete_automation]) + for func in builtin_functions: callable = get_async_tool_function_and_apply_extra_params( func, diff --git a/src/lib/components/workspace/Models/BuiltinTools.svelte b/src/lib/components/workspace/Models/BuiltinTools.svelte index cc94cddb33..b171c85e46 100644 --- a/src/lib/components/workspace/Models/BuiltinTools.svelte +++ b/src/lib/components/workspace/Models/BuiltinTools.svelte @@ -46,6 +46,10 @@ tasks: { label: $i18n.t('Task Management'), description: $i18n.t('Break down complex requests into trackable steps') + }, + automations: { + label: $i18n.t('Automations'), + description: $i18n.t('Create and manage scheduled automations') } }; From 09f6d7ba57d2aaad83ad0d29d005feb7157776a1 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Sat, 11 Apr 2026 16:55:20 -0600 Subject: [PATCH 16/67] refac --- backend/open_webui/models/automations.py | 36 ++++++++++++++++++++++- backend/open_webui/routers/automations.py | 13 +++++++- src/routes/(app)/automations/+page.svelte | 10 ++++--- 3 files changed, 53 insertions(+), 6 deletions(-) diff --git a/backend/open_webui/models/automations.py b/backend/open_webui/models/automations.py index 485f097d5f..e6b8421539 100644 --- a/backend/open_webui/models/automations.py +++ b/backend/open_webui/models/automations.py @@ -45,7 +45,10 @@ class AutomationRun(Base): error = Column(Text, nullable=True) created_at = Column(BigInteger, nullable=False) - __table_args__ = (Index('ix_automation_run_automation_id', 'automation_id'),) + __table_args__ = ( + Index('ix_automation_run_automation_id', 'automation_id'), + Index('ix_automation_run_aid_created', 'automation_id', 'created_at'), + ) #################### @@ -308,6 +311,37 @@ class AutomationRunTable: ) return AutomationRunModel.model_validate(row) if row else None + def get_latest_batch( + self, automation_ids: list[str], db: Optional[Session] = None + ) -> dict[str, AutomationRunModel]: + """Fetch the latest run for each automation in a single query.""" + if not automation_ids: + return {} + with get_db_context(db) as db: + # Subquery: max created_at per automation_id + subq = ( + db.query( + AutomationRun.automation_id, + func.max(AutomationRun.created_at).label('max_created'), + ) + .filter(AutomationRun.automation_id.in_(automation_ids)) + .group_by(AutomationRun.automation_id) + .subquery() + ) + rows = ( + db.query(AutomationRun) + .join( + subq, + (AutomationRun.automation_id == subq.c.automation_id) + & (AutomationRun.created_at == subq.c.max_created), + ) + .all() + ) + return { + row.automation_id: AutomationRunModel.model_validate(row) + for row in rows + } + def get_by_automation( self, automation_id: str, diff --git a/backend/open_webui/routers/automations.py b/backend/open_webui/routers/automations.py index 803f59a6f2..f4608d0770 100644 --- a/backend/open_webui/routers/automations.py +++ b/backend/open_webui/routers/automations.py @@ -61,6 +61,7 @@ def check_automation_access(automation, user): def enrich_automation(automation: AutomationModel, db: Session, tz: str = None) -> AutomationResponse: + """Full enrichment for single-item views (includes next_runs computation).""" last_run = AutomationRuns.get_latest(automation.id, db=db) return AutomationResponse( **automation.model_dump(), @@ -97,8 +98,18 @@ async def get_automation_items( db=db, ) + # Batch-fetch latest runs in a single query instead of N+1 + ids = [item.id for item in result.items] + latest_runs = AutomationRuns.get_latest_batch(ids, db=db) if ids else {} + return { - 'items': [enrich_automation(item, db, tz=user.timezone) for item in result.items], + 'items': [ + AutomationResponse( + **item.model_dump(), + last_run=latest_runs.get(item.id), + ) + for item in result.items + ], 'total': result.total, } diff --git a/src/routes/(app)/automations/+page.svelte b/src/routes/(app)/automations/+page.svelte index 7919ec43ee..e48dd3be37 100644 --- a/src/routes/(app)/automations/+page.svelte +++ b/src/routes/(app)/automations/+page.svelte @@ -47,8 +47,8 @@ let page = 1; - // Debounce only query changes - $: if (query !== undefined) { + // Debounce only query changes (gate behind loaded to prevent double-fetch on mount) + $: if (loaded && query !== undefined) { loading = true; clearTimeout(searchDebounceTimer); searchDebounceTimer = setTimeout(() => { @@ -57,8 +57,8 @@ }, 300); } - // Immediate response to page/filter changes - $: if (page && statusFilter !== undefined) { + // Immediate response to page/filter changes (gate behind loaded) + $: if (loaded && page && statusFilter !== undefined) { getAutomationList(); } @@ -171,6 +171,8 @@ } loaded = true; + // Explicit initial fetch — reactive blocks will handle subsequent changes + await getAutomationList(); return () => { clearTimeout(searchDebounceTimer); From ee9db91df02120e1e3651e8881734966b710ad52 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Sat, 11 Apr 2026 17:06:49 -0600 Subject: [PATCH 17/67] refac --- src/lib/components/chat/Chat.svelte | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/src/lib/components/chat/Chat.svelte b/src/lib/components/chat/Chat.svelte index 5b96f82c2f..a3feb6a269 100644 --- a/src/lib/components/chat/Chat.svelte +++ b/src/lib/components/chat/Chat.svelte @@ -1232,8 +1232,13 @@ desktopEvent.set(null); if (event.type === 'call') { - showCallOverlay.set(true); - showControls.set(true); + // Defer to next macrotask so the call overlay isn't clobbered by + // showControlsSubscribe's initial callback (value=false → set(false)) + // which runs as a pending microtask after this function. + setTimeout(() => { + showCallOverlay.set(true); + showControls.set(true); + }, 0); } else if (event.type === 'query') { const query = event.data?.query; const eventFiles = event.data?.files; From 406251c2f358ffabce4d631c98c6f2c879feae5c Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Sat, 11 Apr 2026 17:06:58 -0600 Subject: [PATCH 18/67] enh: automation --- backend/open_webui/config.py | 12 ++++++++ backend/open_webui/main.py | 4 +++ backend/open_webui/models/automations.py | 4 +++ backend/open_webui/routers/auths.py | 12 ++++++++ backend/open_webui/routers/automations.py | 34 +++++++++++++++++++++++ backend/open_webui/utils/automations.py | 19 +++++++++++++ 6 files changed, 85 insertions(+) diff --git a/backend/open_webui/config.py b/backend/open_webui/config.py index e0e853e3ab..b599ff575d 100644 --- a/backend/open_webui/config.py +++ b/backend/open_webui/config.py @@ -1541,6 +1541,18 @@ ENABLE_CHANNELS = PersistentConfig( os.environ.get('ENABLE_CHANNELS', 'False').lower() == 'true', ) +AUTOMATION_MAX_COUNT = PersistentConfig( + 'AUTOMATION_MAX_COUNT', + 'automations.max_count', + os.environ.get('AUTOMATION_MAX_COUNT', ''), +) + +AUTOMATION_MIN_INTERVAL = PersistentConfig( + 'AUTOMATION_MIN_INTERVAL', + 'automations.min_interval', + os.environ.get('AUTOMATION_MIN_INTERVAL', ''), +) + ENABLE_NOTES = PersistentConfig( 'ENABLE_NOTES', 'notes.enable', diff --git a/backend/open_webui/main.py b/backend/open_webui/main.py index 03bb651089..8fc56a6659 100644 --- a/backend/open_webui/main.py +++ b/backend/open_webui/main.py @@ -383,6 +383,8 @@ from open_webui.config import ( API_KEYS_ALLOWED_ENDPOINTS, ENABLE_FOLDERS, FOLDER_MAX_FILE_COUNT, + AUTOMATION_MAX_COUNT, + AUTOMATION_MIN_INTERVAL, ENABLE_CHANNELS, ENABLE_NOTES, ENABLE_USER_STATUS, @@ -874,6 +876,8 @@ app.state.config.BANNERS = WEBUI_BANNERS app.state.config.ENABLE_FOLDERS = ENABLE_FOLDERS app.state.config.FOLDER_MAX_FILE_COUNT = FOLDER_MAX_FILE_COUNT +app.state.config.AUTOMATION_MAX_COUNT = AUTOMATION_MAX_COUNT +app.state.config.AUTOMATION_MIN_INTERVAL = AUTOMATION_MIN_INTERVAL app.state.config.ENABLE_CHANNELS = ENABLE_CHANNELS app.state.config.ENABLE_NOTES = ENABLE_NOTES app.state.config.ENABLE_COMMUNITY_SHARING = ENABLE_COMMUNITY_SHARING diff --git a/backend/open_webui/models/automations.py b/backend/open_webui/models/automations.py index e6b8421539..7a6bccb9a0 100644 --- a/backend/open_webui/models/automations.py +++ b/backend/open_webui/models/automations.py @@ -143,6 +143,10 @@ class AutomationTable: db.refresh(row) return AutomationModel.model_validate(row) + def count_by_user(self, user_id: str, db: Optional[Session] = None) -> int: + with get_db_context(db) as db: + return db.query(Automation).filter_by(user_id=user_id).count() + def get_by_id(self, id: str, db: Optional[Session] = None) -> Optional[AutomationModel]: with get_db_context(db) as db: row = db.get(Automation, id) diff --git a/backend/open_webui/routers/auths.py b/backend/open_webui/routers/auths.py index 8cca63baf0..5e054596c1 100644 --- a/backend/open_webui/routers/auths.py +++ b/backend/open_webui/routers/auths.py @@ -950,6 +950,8 @@ async def get_admin_config(request: Request, user=Depends(get_admin_user)): 'ENABLE_MESSAGE_RATING': request.app.state.config.ENABLE_MESSAGE_RATING, 'ENABLE_FOLDERS': request.app.state.config.ENABLE_FOLDERS, 'FOLDER_MAX_FILE_COUNT': request.app.state.config.FOLDER_MAX_FILE_COUNT, + 'AUTOMATION_MAX_COUNT': request.app.state.config.AUTOMATION_MAX_COUNT, + 'AUTOMATION_MIN_INTERVAL': request.app.state.config.AUTOMATION_MIN_INTERVAL, 'ENABLE_CHANNELS': request.app.state.config.ENABLE_CHANNELS, 'ENABLE_MEMORIES': request.app.state.config.ENABLE_MEMORIES, 'ENABLE_NOTES': request.app.state.config.ENABLE_NOTES, @@ -976,6 +978,8 @@ class AdminConfig(BaseModel): ENABLE_MESSAGE_RATING: bool ENABLE_FOLDERS: bool FOLDER_MAX_FILE_COUNT: Optional[int | str] = None + AUTOMATION_MAX_COUNT: Optional[int | str] = None + AUTOMATION_MIN_INTERVAL: Optional[int | str] = None ENABLE_CHANNELS: bool ENABLE_MEMORIES: bool ENABLE_NOTES: bool @@ -1001,6 +1005,12 @@ async def update_admin_config(request: Request, form_data: AdminConfig, user=Dep request.app.state.config.FOLDER_MAX_FILE_COUNT = ( int(form_data.FOLDER_MAX_FILE_COUNT) if form_data.FOLDER_MAX_FILE_COUNT else '' ) + request.app.state.config.AUTOMATION_MAX_COUNT = ( + int(form_data.AUTOMATION_MAX_COUNT) if form_data.AUTOMATION_MAX_COUNT else '' + ) + request.app.state.config.AUTOMATION_MIN_INTERVAL = ( + int(form_data.AUTOMATION_MIN_INTERVAL) if form_data.AUTOMATION_MIN_INTERVAL else '' + ) request.app.state.config.ENABLE_CHANNELS = form_data.ENABLE_CHANNELS request.app.state.config.ENABLE_MEMORIES = form_data.ENABLE_MEMORIES request.app.state.config.ENABLE_NOTES = form_data.ENABLE_NOTES @@ -1042,6 +1052,8 @@ async def update_admin_config(request: Request, form_data: AdminConfig, user=Dep 'ENABLE_MESSAGE_RATING': request.app.state.config.ENABLE_MESSAGE_RATING, 'ENABLE_FOLDERS': request.app.state.config.ENABLE_FOLDERS, 'FOLDER_MAX_FILE_COUNT': request.app.state.config.FOLDER_MAX_FILE_COUNT, + 'AUTOMATION_MAX_COUNT': request.app.state.config.AUTOMATION_MAX_COUNT, + 'AUTOMATION_MIN_INTERVAL': request.app.state.config.AUTOMATION_MIN_INTERVAL, 'ENABLE_CHANNELS': request.app.state.config.ENABLE_CHANNELS, 'ENABLE_MEMORIES': request.app.state.config.ENABLE_MEMORIES, 'ENABLE_NOTES': request.app.state.config.ENABLE_NOTES, diff --git a/backend/open_webui/routers/automations.py b/backend/open_webui/routers/automations.py index f4608d0770..0f85115720 100644 --- a/backend/open_webui/routers/automations.py +++ b/backend/open_webui/routers/automations.py @@ -19,6 +19,7 @@ from open_webui.utils.automations import ( next_run_ns, next_n_runs_ns, execute_automation, + rrule_interval_seconds, ) from open_webui.utils.auth import get_verified_user, get_admin_user from open_webui.utils.access_control import has_permission @@ -60,6 +61,35 @@ def check_automation_access(automation, user): ) +def check_automation_limits(request, user, rrule_str: str, db, is_create: bool = False): + """Enforce global automation limits. Admins bypass all checks.""" + if user.role == 'admin': + return + + # Max count (create only) + if is_create: + max_count = request.app.state.config.AUTOMATION_MAX_COUNT + if max_count: + max_count = int(max_count) + if max_count > 0 and Automations.count_by_user(user.id, db=db) >= max_count: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail=f'Automation limit reached ({max_count})', + ) + + # Min interval (create + update) + min_interval = request.app.state.config.AUTOMATION_MIN_INTERVAL + if min_interval: + min_interval = int(min_interval) + if min_interval > 0: + interval = rrule_interval_seconds(rrule_str) + if interval is not None and interval < min_interval: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=f'Schedule too frequent. Minimum interval is {min_interval} seconds.', + ) + + def enrich_automation(automation: AutomationModel, db: Session, tz: str = None) -> AutomationResponse: """Full enrichment for single-item views (includes next_runs computation).""" last_run = AutomationRuns.get_latest(automation.id, db=db) @@ -135,6 +165,8 @@ async def create_new_automation( detail=str(e), ) + check_automation_limits(request, user, form_data.data.rrule, db, is_create=True) + # Validate terminal server exists if linked if form_data.data.terminal and form_data.data.terminal.server_id: connections = request.app.state.config.TERMINAL_SERVER_CONNECTIONS or [] @@ -192,6 +224,8 @@ async def update_automation_by_id( detail=str(e), ) + check_automation_limits(request, user, form_data.data.rrule, db, is_create=False) + # Validate terminal server exists if linked if form_data.data.terminal and form_data.data.terminal.server_id: connections = request.app.state.config.TERMINAL_SERVER_CONNECTIONS or [] diff --git a/backend/open_webui/utils/automations.py b/backend/open_webui/utils/automations.py index 4d9eb2fb6c..5d28307f7a 100644 --- a/backend/open_webui/utils/automations.py +++ b/backend/open_webui/utils/automations.py @@ -92,6 +92,25 @@ def next_n_runs_ns(s: str, n: int = 5, tz: str = None) -> list[int]: return result +def rrule_interval_seconds(s: str) -> Optional[int]: + """Approximate interval between recurrences in seconds. + + Returns None for one-shot (COUNT=1) schedules or rules + with fewer than two future occurrences. + """ + if 'COUNT=1' in s: + return None + rule = _parse_rule(s) + now = datetime.now() + first = rule.after(now) + if first is None: + return None + second = rule.after(first) + if second is None: + return None + return int((second - first).total_seconds()) + + ############################ # Worker Loop ############################ From 92dfa3f2f28922a6adc2f88433f9259644443593 Mon Sep 17 00:00:00 2001 From: G30 <50341825+silentoplayz@users.noreply.github.com> Date: Sat, 11 Apr 2026 19:10:11 -0400 Subject: [PATCH 19/67] fix(backend): provide fallback strings for webhook UserNameResponse to prevent Pydantic validation error (#23414) --- backend/open_webui/routers/channels.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/backend/open_webui/routers/channels.py b/backend/open_webui/routers/channels.py index 5241fe7445..714bde8a85 100644 --- a/backend/open_webui/routers/channels.py +++ b/backend/open_webui/routers/channels.py @@ -788,8 +788,8 @@ async def get_pinned_channel_messages( webhook_info = message.meta.get('webhook') if message.meta else None if webhook_info: user_info = UserNameResponse( - id=webhook_info.get('id'), - name=webhook_info.get('name'), + id=webhook_info.get('id') or '', + name=webhook_info.get('name') or 'Webhook', role='webhook', ) elif message.user_id in users: From b0df5272243617a765e65f4876d9fda6671c85fd Mon Sep 17 00:00:00 2001 From: Toru Suzuki <35589743+zolgear@users.noreply.github.com> Date: Mon, 13 Apr 2026 01:04:35 +0900 Subject: [PATCH 20/67] i18n: Update Japanese translation (#23617) --- src/lib/i18n/locales/ja-JP/translation.json | 70 ++++++++++----------- 1 file changed, 35 insertions(+), 35 deletions(-) diff --git a/src/lib/i18n/locales/ja-JP/translation.json b/src/lib/i18n/locales/ja-JP/translation.json index c4032c2728..b559404dc9 100644 --- a/src/lib/i18n/locales/ja-JP/translation.json +++ b/src/lib/i18n/locales/ja-JP/translation.json @@ -11,7 +11,7 @@ "{{ models }}": "{{ モデル }}", "{{COUNT}} Available Tools": "{{COUNT}} 個の有効なツール", "{{COUNT}} characters": "{{COUNT}} 文字", - "{{COUNT}} extracted lines": "", + "{{COUNT}} extracted lines": "{{COUNT}} 行を抽出", "{{COUNT}} files": "", "{{COUNT}} hidden lines": "{{COUNT}} 行が非表示", "{{COUNT}} members": "{{COUNT}} メンバー", @@ -185,7 +185,7 @@ "Are you sure you want to delete all chats? This action cannot be undone.": "すべてのチャットを削除しますか? この操作は元に戻すことができません。", "Are you sure you want to delete this channel?": "このチャンネルを削除しますか?", "Are you sure you want to delete this connection? This action cannot be undone.": "", - "Are you sure you want to delete this memory? This action cannot be undone.": "", + "Are you sure you want to delete this memory? This action cannot be undone.": "このメモリをクリアしますか? この操作は元に戻すことができません。", "Are you sure you want to delete this message?": "このメッセージを削除しますか?", "Are you sure you want to delete this version? Child versions will be relinked to this version's parent.": "", "Are you sure you want to delete this?": "", @@ -198,7 +198,7 @@ "Assistant": "アシスタント", "Async Embedding Processing": "", "Attach File From Knowledge": "ナレッジからファイルを添付", - "Attach Files": "", + "Attach Files": "ファイルを追加", "Attach Knowledge": "ナレッジを追加", "Attach Notes": "ノートを追加", "Attach Webpage": "ウェブページを追加", @@ -221,13 +221,13 @@ "AUTOMATIC1111 Base URL": "AUTOMATIC1111 ベース URL", "AUTOMATIC1111 Base URL is required.": "AUTOMATIC1111 ベース URL が必要です。", "Automatically inject system tools in native function calling mode (e.g., timestamps, memory, chat history, notes, etc.)": "ネイティブの関数呼び出しモードにおいて、システムツール(例: タイムスタンプ、メモリー、チャット履歴、ノートなど)を自動的に注入します", - "Automation": "", - "Automation created": "", - "Automation Name": "", - "Automation title": "", - "Automation triggered": "", - "Automation updated": "", - "Automations": "", + "Automation": "オートメーション", + "Automation created": "オートメーションを作成しました", + "Automation Name": "オートメーション名", + "Automation title": "オートメーションのタイトル", + "Automation triggered": "オートメーションが実行されました", + "Automation updated": "オートメーションを更新しました", + "Automations": "オートメーション", "Available list": "利用可能リスト", "Available models": "", "Available Tools": "利用可能ツール", @@ -392,7 +392,7 @@ "Concurrent Requests": "同時リクエスト", "Config": "", "Config imported successfully": "設定のインポートに成功しました", - "Configuration": "", + "Configuration": "設定", "Configure": "設定", "Confirm": "確認", "Confirm Password": "パスワードの確認", @@ -462,7 +462,7 @@ "Create new secret key": "新しいシークレットキーを作成", "Create note": "ノートを作成", "Create Note": "ノートを作成", - "Create scheduled prompts that run automatically on a recurring basis.": "", + "Create scheduled prompts that run automatically on a recurring basis.": "定期的に自動実行されるプロンプトを作成します。", "Create your first note by clicking on the plus button below.": "プラスボタンをクリックして最初のノートを作成します。", "Created at": "作成日時", "Created At": "作成日時", @@ -515,13 +515,13 @@ "Delete All": "すべて削除する", "Delete All Chats": "すべてのチャットを削除", "Delete all contents inside this folder": "", - "Delete automation?": "", + "Delete automation?": "オートメーションを削除しますか?", "Delete Chat": "チャットを削除", "Delete chat?": "チャットを削除しますか?", "Delete File": "", "Delete folder?": "フォルダーを削除しますか?", "Delete function?": "Functionを削除しますか?", - "Delete Memory?": "", + "Delete Memory?": "メモリを削除しますか?", "Delete Message": "メッセージを削除", "Delete message?": "メッセージを削除しますか?", "Delete Model": "", @@ -752,7 +752,7 @@ "Enter Perplexity Search API URL": "", "Enter Playwright Timeout": "Playwrightタイムアウトを入力", "Enter Playwright WebSocket URL": "Playwright WebSocket URLを入力", - "Enter prompt here.": "", + "Enter prompt here.": "ここにプロンプトを入力", "Enter proxy URL (e.g. https://user:password@host:port)": "プロキシURLを入力 (例: https://user:password@host:port)", "Enter reasoning effort": "推論の努力を入力", "Enter Score": "スコアを入力", @@ -777,7 +777,7 @@ "Enter system prompt here": "システムプロンプトをここに入力", "Enter Tavily API Key": "Tavily API Keyを入力", "Enter Tavily Extract Depth": "Tavily Extract Depthを入力", - "Enter the prompt instructions for this automation...": "", + "Enter the prompt instructions for this automation...": "このオートメーションのプロンプトを入力してください...", "Enter the public URL of your WebUI. This URL will be used to generate links in the notifications.": "WebUIの公開URLを入力してください。このURLは通知でリンクを生成するために使用されます。", "Enter the URL of the function to import": "インポートするFunctionのURLを入力", "Enter the URL to import": "インポートするURLを入力", @@ -835,7 +835,7 @@ "Execute code": "", "Execute code for analysis": "コードの分析に実行", "Executing **{{NAME}}**...": "**{{NAME}}**を実行中...", - "Execution Logs": "", + "Execution Logs": "実行ログ", "Expand": "展開", "Experimental": "実験的", "Explain": "説明", @@ -964,7 +964,7 @@ "Form": "フォーム", "Format Lines": "出力テキストをフォーマット", "Format the lines in the output. Defaults to False. If set to True, the lines will be formatted to detect inline math and styles.": "出力をフォーマットする。デフォルトでは無効です。有効にすると、インライン数式やスタイルを検出しフォーマットします。", - "Formatting may be inconsistent from source.": "", + "Formatting may be inconsistent from source.": "元のデータにより、書式が一致しない場合があります。", "Forward": "", "Forwards system user OAuth access token to authenticate": "", "Forwards system user session credentials to authenticate": "システムユーザーセッションの資格情報を転送して認証する", @@ -1101,7 +1101,7 @@ "Insert Suggestion Prompt to Input": "", "Install from Github URL": "Github URLからインストール", "Instant Auto-Send After Voice Transcription": "音声文字変換後に自動送信", - "Instructions": "", + "Instructions": "指示", "Integration": "連携", "Integrations": "連携", "Interface": "インターフェース", @@ -1161,7 +1161,7 @@ "Last 90 days": "", "Last Active": "最終アクティブ", "Last Modified": "最終変更", - "Last ran": "", + "Last ran": "最終実行", "Last reply": "最終応答", "LDAP": "LDAP", "LDAP server updated": "LDAPサーバーの更新に成功しました", @@ -1318,11 +1318,11 @@ "Name": "名前", "Name and ID are required, please fill them out": "名前とIDは必須です。項目を入力してください。", "Name your knowledge base": "ナレッジベースに名前を付ける", - "Name, prompt, and model are required": "", + "Name, prompt, and model are required": "名前、プロンプト、モデルの選択が必要です", "Native": "ネイティブ", - "Never": "", + "Never": "なし", "New": "", - "New Automation": "", + "New Automation": "新しいオートメーション", "New Button": "新しいボタン", "New Chat": "新しいチャット", "New File": "", @@ -1341,11 +1341,11 @@ "New Webhook": "", "new-channel": "新しいチャンネル", "Next message": "次のメッセージ", - "Next run": "", + "Next run": "次回実行", "No access grants. Private to you.": "アクセス権は付与されていません。あなただけが利用できます。", "No activity data": "", "No authentication": "", - "No automations found": "", + "No automations found": "オートメーションが見つかりません", "No chats found": "チャットが見つかりません。", "No chats found for this user.": "このユーザーのチャットが見つかりません。", "No chats found.": "チャットが見つかりません。", @@ -1493,7 +1493,7 @@ "Password": "パスワード", "Passwords do not match.": "パスワードが一致しません。", "Paste Large Text as File": "大きなテキストをファイルとして貼り付ける", - "Paused": "", + "Paused": "停止中", "PDF document (.pdf)": "PDF ドキュメント (.pdf)", "PDF Extract Images (OCR)": "PDF 画像抽出 (OCR)", "PDF Loader Mode": "", @@ -1634,7 +1634,7 @@ "Renamed to {{name}}": "", "Render Markdown in Previews": "", "Reorder Models": "モデルを並べ替え", - "Repeats": "", + "Repeats": "繰り返し", "Reply": "", "Reply in Thread": "スレッドで返信", "Reply to thread...": "", @@ -1666,8 +1666,8 @@ "RTL": "RTL", "Run": "実行", "Run All": "", - "Run now": "", - "Run Now": "", + "Run now": "今すぐ実行", + "Run Now": "今すぐ実行", "Running": "実行中", "Running...": "実行中...", "Runs embedding tasks concurrently to speed up processing. Turn off if rate limits become an issue.": "", @@ -1678,7 +1678,7 @@ "Save Chat": "チャットを保存", "Saved": "保存しました。", "Saving chat logs directly to your browser's storage is no longer supported. Please take a moment to download and delete your chat logs by clicking the button below. Don't worry, you can easily re-import your chat logs to the backend through": "チャットログをブラウザのストレージに直接保存する機能はサポートされなくなりました。下のボタンをクリックして、チャットログをダウンロードして削除してください。ご心配なく。チャットログは、次の方法でバックエンドに簡単に再インポートできます。", - "Schedule": "", + "Schedule": "スケジュール", "Scheduled time must be in the future": "", "Scroll On Branch Change": "ブランチ変更時にスクロール", "Search": "検索", @@ -1686,7 +1686,7 @@ "Search all emojis": "絵文字を検索", "Search and manage user memories": "", "Search and view user chat history": "", - "Search Automations": "", + "Search Automations": "オートメーションを検索", "Search Base": "ベースを検索", "Search channels and channel messages": "", "Search Chats": "チャットの検索", @@ -1702,7 +1702,7 @@ "Search Groups": "グループの検索", "Search In Models": "モデルを検索", "Search Knowledge": "ナレッジベースの検索", - "Search Memories": "", + "Search Memories": "メモリの検索", "Search Models": "モデル検索", "Search Notes": "ノートを検索", "Search options": "検索オプション", @@ -1756,7 +1756,7 @@ "Select how to split message text for TTS requests": "TTSリクエストのテキスト分割方法を選択", "Select Knowledge": "ナレッジベースの選択", "Select Method": "", - "Select model": "", + "Select model": "モデルを選択", "Select only one model to call": "1つのモデルを呼び出すには、1つのモデルを選択してください。", "Select view": "", "Selected model: {{modelName}}": "", @@ -1866,7 +1866,7 @@ "Start of the channel": "チャンネルの開始", "Start Tag": "", "Starting kernel...": "", - "State": "", + "State": "状態", "Status": "ステータス", "Status cleared successfully": "正常にステータスをクリアしました", "Status updated successfully": "正常にステータスを更新しました", @@ -2001,7 +2001,7 @@ "To access the WebUI, please reach out to the administrator. Admins can manage user statuses from the Admin Panel.": "WebUIにアクセスするには、管理者にお問い合わせください。管理者は管理者パネルからユーザーのステータスを管理できます。", "To attach knowledge base here, add them to the \"Knowledge\" workspace first.": "ここにナレッジベースを追加するには、まず \"Knowledge\" ワークスペースに追加してください。", "To learn more about available endpoints, visit our documentation.": "利用可能なエンドポイントについては、ドキュメントを参照してください。", - "To select skills here, add them to the \"Skills\" workspace first.": "", + "To select skills here, add them to the \"Skills\" workspace first.": "ここでSkillを選択するには、まず\"Skills\" ワークスペースに追加してください。", "To select toolkits here, add them to the \"Tools\" workspace first.": "ここでツールキットを選択するには、まず \"Tools\" ワークスペースに追加してください。", "Toast notifications for new updates": "新しい更新のトースト通知", "Today": "今日", From e790e7be7abfa2740bc16c748351dcc6e9d62699 Mon Sep 17 00:00:00 2001 From: Classic298 <27028174+Classic298@users.noreply.github.com> Date: Sun, 12 Apr 2026 18:06:33 +0200 Subject: [PATCH 21/67] fix: enforce model access control on /responses endpoint (#23481) The /responses proxy endpoint only required authentication via get_verified_user but did not check per-model access grants. This allowed any authenticated user to access any model through this endpoint, bypassing the access control system. Extract a shared check_model_access helper into utils/access_control and replace all inline access control blocks across openai.py and ollama.py (7 locations) with calls to this helper. This eliminates code duplication and prevents future policy drift between endpoints. CWE-862: Missing Authorization CVSS:3.1/AV:N/AC:L/PR:L/UI:N/S:U/C:L/I:N/A:H (6.5 Medium) --- backend/open_webui/routers/ollama.py | 99 ++----------------- backend/open_webui/routers/openai.py | 34 ++----- .../utils/access_control/__init__.py | 44 +++++++++ 3 files changed, 64 insertions(+), 113 deletions(-) diff --git a/backend/open_webui/routers/ollama.py b/backend/open_webui/routers/ollama.py index 93745440c4..728d09e524 100644 --- a/backend/open_webui/routers/ollama.py +++ b/backend/open_webui/routers/ollama.py @@ -47,6 +47,7 @@ from open_webui.internal.db import get_session from open_webui.models.models import Models from open_webui.models.access_grants import AccessGrants from open_webui.models.groups import Groups +from open_webui.utils.access_control import check_model_access from open_webui.utils.misc import ( calculate_sha256, cleanup_response, @@ -1097,29 +1098,9 @@ async def generate_chat_completion( if not bypass_system_prompt: payload = apply_system_prompt_to_body(system, payload, metadata, user) - # Check if user has access to the model - if not bypass_filter and user.role == 'user': - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)} - if not ( - user.id == model_info.user_id - or AccessGrants.has_access( - user_id=user.id, - resource_type='model', - resource_id=model_info.id, - permission='read', - user_group_ids=user_group_ids, - ) - ): - raise HTTPException( - status_code=403, - detail='Model not found', - ) - elif not bypass_filter: - if user.role != 'admin': - raise HTTPException( - status_code=403, - detail='Model not found', - ) + check_model_access(user, model_info, bypass_filter) + else: + check_model_access(user, None, bypass_filter) url, url_idx = await get_ollama_url(request, payload['model'], url_idx) api_config = request.app.state.config.OLLAMA_API_CONFIGS.get( @@ -1206,29 +1187,9 @@ async def generate_openai_completion( if params: payload = apply_model_params_to_body_openai(params, payload) - # Check if user has access to the model - if user.role == 'user': - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)} - if not ( - user.id == model_info.user_id - or AccessGrants.has_access( - user_id=user.id, - resource_type='model', - resource_id=model_info.id, - permission='read', - user_group_ids=user_group_ids, - ) - ): - raise HTTPException( - status_code=403, - detail='Model not found', - ) + check_model_access(user, model_info) else: - if user.role != 'admin': - raise HTTPException( - status_code=403, - detail='Model not found', - ) + check_model_access(user, None) url, url_idx = await get_ollama_url(request, payload['model'], url_idx) api_config = request.app.state.config.OLLAMA_API_CONFIGS.get( @@ -1292,29 +1253,9 @@ async def generate_openai_chat_completion( payload = apply_model_params_to_body_openai(params, payload) payload = apply_system_prompt_to_body(system, payload, metadata, user) - # Check if user has access to the model - if user.role == 'user': - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)} - if not ( - user.id == model_info.user_id - or AccessGrants.has_access( - user_id=user.id, - resource_type='model', - resource_id=model_info.id, - permission='read', - user_group_ids=user_group_ids, - ) - ): - raise HTTPException( - status_code=403, - detail='Model not found', - ) + check_model_access(user, model_info) else: - if user.role != 'admin': - raise HTTPException( - status_code=403, - detail='Model not found', - ) + check_model_access(user, None) url, url_idx = await get_ollama_url(request, payload['model'], url_idx) api_config = request.app.state.config.OLLAMA_API_CONFIGS.get( @@ -1364,29 +1305,9 @@ async def generate_anthropic_messages( if model_info.base_model_id: payload['model'] = model_info.base_model_id - # Check if user has access to the model - if user.role == 'user': - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)} - if not ( - user.id == model_info.user_id - or AccessGrants.has_access( - user_id=user.id, - resource_type='model', - resource_id=model_info.id, - permission='read', - user_group_ids=user_group_ids, - ) - ): - raise HTTPException( - status_code=403, - detail='Model not found', - ) + check_model_access(user, model_info) else: - if user.role != 'admin': - raise HTTPException( - status_code=403, - detail='Model not found', - ) + check_model_access(user, None) url, url_idx = await get_ollama_url(request, payload['model'], url_idx) api_config = request.app.state.config.OLLAMA_API_CONFIGS.get( diff --git a/backend/open_webui/routers/openai.py b/backend/open_webui/routers/openai.py index 836517df9d..83a537f913 100644 --- a/backend/open_webui/routers/openai.py +++ b/backend/open_webui/routers/openai.py @@ -28,6 +28,7 @@ from open_webui.internal.db import get_session from open_webui.models.models import Models from open_webui.models.access_grants import AccessGrants from open_webui.models.groups import Groups +from open_webui.utils.access_control import has_connection_access, check_model_access from open_webui.config import ( CACHE_DIR, ) @@ -1044,29 +1045,9 @@ async def generate_chat_completion( if not bypass_system_prompt: payload = apply_system_prompt_to_body(system, payload, metadata, user) - # Check if user has access to the model - if not bypass_filter and user.role == 'user': - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)} - if not ( - user.id == model_info.user_id - or AccessGrants.has_access( - user_id=user.id, - resource_type='model', - resource_id=model_info.id, - permission='read', - user_group_ids=user_group_ids, - ) - ): - raise HTTPException( - status_code=403, - detail='Model not found', - ) - elif not bypass_filter: - if user.role != 'admin': - raise HTTPException( - status_code=403, - detail='Model not found', - ) + check_model_access(user, model_info, bypass_filter) + else: + check_model_access(user, None, bypass_filter) # Check if model is already in app state cache to avoid expensive get_all_models() call models = request.app.state.OPENAI_MODELS @@ -1340,10 +1321,15 @@ async def responses( Routes to the correct upstream backend based on the model field. """ payload = form_data.model_dump(exclude_none=True) - body = json.dumps(payload) idx = 0 model_id = form_data.model + + # Enforce per-model access control + check_model_access(user, Models.get_model_by_id(model_id), BYPASS_MODEL_ACCESS_CONTROL) + + body = json.dumps(payload) + if model_id: models = request.app.state.OPENAI_MODELS if not models or model_id not in models: diff --git a/backend/open_webui/utils/access_control/__init__.py b/backend/open_webui/utils/access_control/__init__.py index f31c59e158..9c91371384 100644 --- a/backend/open_webui/utils/access_control/__init__.py +++ b/backend/open_webui/utils/access_control/__init__.py @@ -255,3 +255,47 @@ def filter_allowed_access_grants( access_grants = strip_user_access_grants(access_grants) return access_grants + + +def check_model_access( + user: UserModel, + model_info, + bypass_filter: bool = False, +) -> None: + """ + Enforce per-model read access for the given user. + + Raises HTTPException(403) if the user is not authorized. + Does nothing if bypass_filter is True. + + Args: + user: The authenticated user. + model_info: The model record from Models.get_model_by_id(), + or None if the model is not registered. + bypass_filter: If True, skip all access checks (used by + internal callers and BYPASS_MODEL_ACCESS_CONTROL). + """ + from fastapi import HTTPException + + if bypass_filter: + return + + if model_info: + if user.role == 'user': + from open_webui.models.access_grants import AccessGrants + + user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)} + if not ( + user.id == model_info.user_id + or AccessGrants.has_access( + user_id=user.id, + resource_type='model', + resource_id=model_info.id, + permission='read', + user_group_ids=user_group_ids, + ) + ): + raise HTTPException(status_code=403, detail='Model not found') + else: + if user.role != 'admin': + raise HTTPException(status_code=403, detail='Model not found') From 5eab125f13560a7d972c4311152fc68954f4d41b Mon Sep 17 00:00:00 2001 From: Classic298 <27028174+Classic298@users.noreply.github.com> Date: Sun, 12 Apr 2026 18:12:49 +0200 Subject: [PATCH 22/67] fix: sanitize model description HTML with DOMPurify in chat placeholders (#23621) --- src/lib/components/chat/ChatPlaceholder.svelte | 16 +++++++--------- src/lib/components/chat/Placeholder.svelte | 16 +++++++--------- 2 files changed, 14 insertions(+), 18 deletions(-) diff --git a/src/lib/components/chat/ChatPlaceholder.svelte b/src/lib/components/chat/ChatPlaceholder.svelte index 0f44f8693f..ce54ecd551 100644 --- a/src/lib/components/chat/ChatPlaceholder.svelte +++ b/src/lib/components/chat/ChatPlaceholder.svelte @@ -46,11 +46,11 @@ }} > ') - )} + ))} placement="right" > - {@html DOMPurify.sanitize( - marked.parse( - sanitizeResponseContent( - models[selectedModelIdx]?.info?.meta?.description - ).replaceAll('\n', '
') - ) - )} + {@html DOMPurify.sanitize(marked.parse( + sanitizeResponseContent( + models[selectedModelIdx]?.info?.meta?.description + ).replaceAll('\n', '
') + ))} {#if models[selectedModelIdx]?.info?.meta?.user}
diff --git a/src/lib/components/chat/Placeholder.svelte b/src/lib/components/chat/Placeholder.svelte index 1a235b50dd..676ed9846d 100644 --- a/src/lib/components/chat/Placeholder.svelte +++ b/src/lib/components/chat/Placeholder.svelte @@ -165,23 +165,21 @@ {#if models[selectedModelIdx]?.info?.meta?.description ?? null} ') - )} + ))} placement="top" >
- {@html DOMPurify.sanitize( - marked.parse( - sanitizeResponseContent( - models[selectedModelIdx]?.info?.meta?.description ?? '' - ).replaceAll('\n', '
') - ) - )} + {@html DOMPurify.sanitize(marked.parse( + sanitizeResponseContent( + models[selectedModelIdx]?.info?.meta?.description ?? '' + ).replaceAll('\n', '
') + ))}
From 71a39dbac1208e31e885a1e0309d016da47c7f76 Mon Sep 17 00:00:00 2001 From: Classic298 <27028174+Classic298@users.noreply.github.com> Date: Sun, 12 Apr 2026 18:13:51 +0200 Subject: [PATCH 23/67] fix: filter on is_active in channel membership checks (#23623) is_user_channel_member and is_user_channel_manager did not filter on is_active, allowing deactivated members to retain read/write access to group channels via direct API calls. --- backend/open_webui/models/channels.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/backend/open_webui/models/channels.py b/backend/open_webui/models/channels.py index 4d773491d5..10f2e5cad4 100644 --- a/backend/open_webui/models/channels.py +++ b/backend/open_webui/models/channels.py @@ -508,6 +508,7 @@ class ChannelTable: .filter( ChannelMember.channel_id == channel_id, ChannelMember.user_id == user_id, + ChannelMember.is_active.is_(True), ChannelMember.role == 'manager', ) .first() @@ -667,6 +668,7 @@ class ChannelTable: .filter( ChannelMember.channel_id == channel_id, ChannelMember.user_id == user_id, + ChannelMember.is_active.is_(True), ) .first() ) From 36a81ad43b7c0d450079f818a7546eaa517e3d95 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Sun, 12 Apr 2026 11:15:38 -0500 Subject: [PATCH 24/67] refac --- .../workspace/Prompts/PromptEditor.svelte | 38 ++++++++++--------- 1 file changed, 21 insertions(+), 17 deletions(-) diff --git a/src/lib/components/workspace/Prompts/PromptEditor.svelte b/src/lib/components/workspace/Prompts/PromptEditor.svelte index 66db6df208..5a2cf1389b 100644 --- a/src/lib/components/workspace/Prompts/PromptEditor.svelte +++ b/src/lib/components/workspace/Prompts/PromptEditor.svelte @@ -82,23 +82,27 @@ loading = true; if (validateCommandString(command)) { - await onSubmit({ - id: prompt?.id, - name, - command, - content, - tags: tags.map((tag) => tag.name), - access_grants: accessGrants, - commit_message: commitMessage || undefined, - is_production: isProduction - }); - showEditModal = false; - commitMessage = ''; - isProduction = true; - await loadHistory(true); // Reset and reload - // Select the newest version after saving - if (history.length > 0) { - selectedHistoryEntry = history[0]; + try { + await onSubmit({ + id: prompt?.id, + name, + command, + content, + tags: tags.map((tag) => tag.name), + access_grants: accessGrants, + commit_message: commitMessage || undefined, + is_production: isProduction + }); + showEditModal = false; + commitMessage = ''; + isProduction = true; + await loadHistory(true); // Reset and reload + // Select the newest version after saving + if (history.length > 0) { + selectedHistoryEntry = history[0]; + } + } catch (error) { + toast.error(`${error}`); } } else { toast.error( From 96a0b3239b1aadb23fc359bf10849c9ba12fd6ec Mon Sep 17 00:00:00 2001 From: Classic298 <27028174+Classic298@users.noreply.github.com> Date: Sun, 12 Apr 2026 18:28:41 +0200 Subject: [PATCH 25/67] fix: prevent first-user admin race in LDAP and OAuth registration (#23626) Both LDAP and OAuth registration checked user count before insert to determine whether to assign admin role. With multiple workers, concurrent first-user registrations could each see zero users and both create admin accounts. Applies the insert-first-check-after pattern already used by signup_handler: insert with DEFAULT_USER_ROLE, then atomically check get_num_users()==1 and promote only the sole user to admin. --- backend/open_webui/routers/auths.py | 12 +++++++++--- backend/open_webui/utils/oauth.py | 19 ++++++++++++++++--- 2 files changed, 25 insertions(+), 6 deletions(-) diff --git a/backend/open_webui/routers/auths.py b/backend/open_webui/routers/auths.py index 5e054596c1..8d7581dd4b 100644 --- a/backend/open_webui/routers/auths.py +++ b/backend/open_webui/routers/auths.py @@ -479,19 +479,25 @@ async def ldap_auth( user = Users.get_user_by_email(email, db=db) if not user: try: - role = 'admin' if not Users.has_users(db=db) else request.app.state.config.DEFAULT_USER_ROLE - + # Insert with default role first to avoid TOCTOU race on + # first-user registration. Matches signup_handler pattern. user = Auths.insert_new_auth( email=email, password=str(uuid.uuid4()), name=cn, - role=role, + role=request.app.state.config.DEFAULT_USER_ROLE, db=db, ) if not user: raise HTTPException(500, detail=ERROR_MESSAGES.CREATE_USER_ERROR) + # Atomically check if this is the only user *after* the + # insert. Only the single user present should become admin. + if Users.get_num_users(db=db) == 1: + Users.update_user_role_by_id(user.id, 'admin', db=db) + user = Users.get_user_by_id(user.id, db=db) + apply_default_group_assignment( request.app.state.config.DEFAULT_GROUP_ID, user.id, diff --git a/backend/open_webui/utils/oauth.py b/backend/open_webui/utils/oauth.py index 8aaadc271a..100df7a219 100644 --- a/backend/open_webui/utils/oauth.py +++ b/backend/open_webui/utils/oauth.py @@ -1109,9 +1109,12 @@ class OAuthManager: log.debug('Assigning the only user the admin role') return 'admin' if not user and user_count == 0: - # If there are no users, assign the role "admin", as the first user will be an admin - log.debug('Assigning the first user the admin role') - return 'admin' + # First-user bootstrap: skip role management gating so the + # instance can be initialized. We intentionally return the + # default role here (not 'admin') — admin promotion happens + # race-safely *after* insert via get_num_users() == 1. + log.debug('First user bootstrap: using default role (admin promotion deferred to post-insert)') + return auth_manager_config.DEFAULT_USER_ROLE if auth_manager_config.ENABLE_OAUTH_ROLE_MANAGEMENT: log.debug('Running OAUTH Role management') @@ -1577,6 +1580,16 @@ class OAuthManager: db=db, ) + if not user: + raise HTTPException(500, detail=ERROR_MESSAGES.CREATE_USER_ERROR) + + # Atomically check if this is the only user *after* the + # insert to avoid TOCTOU race on first-user registration. + # Matches signup_handler pattern. + if Users.get_num_users(db=db) == 1: + Users.update_user_role_by_id(user.id, 'admin', db=db) + user = Users.get_user_by_id(user.id, db=db) + if auth_manager_config.WEBHOOK_URL: await post_webhook( WEBUI_NAME, From d3df8f1f372411314be9121fbf61d107939fa258 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Sun, 12 Apr 2026 12:28:38 -0500 Subject: [PATCH 26/67] refac --- backend/open_webui/utils/tools.py | 15 +++++++++++++-- 1 file changed, 13 insertions(+), 2 deletions(-) diff --git a/backend/open_webui/utils/tools.py b/backend/open_webui/utils/tools.py index aff33caa4d..2223f202a0 100644 --- a/backend/open_webui/utils/tools.py +++ b/backend/open_webui/utils/tools.py @@ -734,20 +734,31 @@ def get_tool_specs(tool_module: object) -> list[dict]: return specs -def resolve_schema(schema, components): +def resolve_schema(schema, components, resolved_schemas=None): """ Recursively resolves a JSON schema using OpenAPI components. """ if not schema: return {} + if resolved_schemas is None: + resolved_schemas = set() + if '$ref' in schema: ref_path = schema['$ref'] + schema_name = ref_path.split('/')[-1] + + if schema_name in resolved_schemas: + # Avoid infinite recursion on circular references + return {} + + resolved_schemas.add(schema_name) + ref_parts = ref_path.strip('#/').split('/') resolved = components for part in ref_parts[1:]: # Skip the initial 'components' resolved = resolved.get(part, {}) - return resolve_schema(resolved, components) + return resolve_schema(resolved, components, resolved_schemas) resolved_schema = copy.deepcopy(schema) From 4498c21f4cd48c7e42d99b16af947b20273ee8dd Mon Sep 17 00:00:00 2001 From: Classic298 <27028174+Classic298@users.noreply.github.com> Date: Sun, 12 Apr 2026 19:29:26 +0200 Subject: [PATCH 27/67] fix: enforce model access control on Ollama generate, show, embed, embeddings endpoints (#23631) These four endpoints checked model existence but never verified the user has read access via AccessGrants, allowing any authenticated user to use restricted models. Uses the canonical check_model_access helper from utils.access_control. --- backend/open_webui/routers/ollama.py | 17 +++++++++++++++-- 1 file changed, 15 insertions(+), 2 deletions(-) diff --git a/backend/open_webui/routers/ollama.py b/backend/open_webui/routers/ollama.py index 728d09e524..7d5a316fbd 100644 --- a/backend/open_webui/routers/ollama.py +++ b/backend/open_webui/routers/ollama.py @@ -797,11 +797,14 @@ async def show_model_info(request: Request, form_data: ModelNameForm, user=Depen form_data = form_data.model_dump(exclude_none=True) form_data['model'] = form_data.get('model', form_data.get('name')) + model = form_data.get('model') + + # Enforce per-model access control + check_model_access(user, Models.get_model_by_id(model), BYPASS_MODEL_ACCESS_CONTROL) + await get_all_models(request, user=user) models = request.app.state.OLLAMA_MODELS - model = form_data.get('model') - if model not in models: raise HTTPException( status_code=400, @@ -846,6 +849,9 @@ async def embed( log.info(f'generate_ollama_batch_embeddings {form_data}') + # Enforce per-model access control + check_model_access(user, Models.get_model_by_id(form_data.model), BYPASS_MODEL_ACCESS_CONTROL) + if url_idx is None: model = form_data.model @@ -902,6 +908,9 @@ async def embeddings( log.info(f'generate_ollama_embeddings {form_data}') + # Enforce per-model access control + check_model_access(user, Models.get_model_by_id(form_data.model), BYPASS_MODEL_ACCESS_CONTROL) + if url_idx is None: model = form_data.model @@ -964,11 +973,15 @@ async def generate_completion( if not request.app.state.config.ENABLE_OLLAMA_API: raise HTTPException(status_code=503, detail='Ollama API is disabled') + # Enforce per-model access control + check_model_access(user, Models.get_model_by_id(form_data.model), BYPASS_MODEL_ACCESS_CONTROL) + if url_idx is None: await get_all_models(request, user=user) models = request.app.state.OLLAMA_MODELS model = form_data.model + if model in models: url_idx = random.choice(models[model]['urls']) else: From a2a9a3a42a5be5723e33cae6363a4f5fa8005f43 Mon Sep 17 00:00:00 2001 From: Classic298 <27028174+Classic298@users.noreply.github.com> Date: Sun, 12 Apr 2026 19:29:45 +0200 Subject: [PATCH 28/67] fix: prevent path traversal via model name in Azure deployment URLs (#23629) The model name from user input was interpolated directly into Azure deployment URL paths without validation. A user could send a model name like '../../management/foo' to traverse the URL path and hit unintended Azure endpoints with the admin's API key. Adds _sanitize_model_for_url that rejects path separators and traversal sequences, and percent-encodes the name. Applied at convert_to_azure_payload (covers chat completions + proxy) and the responses endpoint's direct URL construction. --- backend/open_webui/routers/openai.py | 26 ++++++++++++++++++++++++-- 1 file changed, 24 insertions(+), 2 deletions(-) diff --git a/backend/open_webui/routers/openai.py b/backend/open_webui/routers/openai.py index 83a537f913..14225a243f 100644 --- a/backend/open_webui/routers/openai.py +++ b/backend/open_webui/routers/openai.py @@ -4,7 +4,7 @@ import json import logging import re from typing import Optional -from urllib.parse import urlparse +from urllib.parse import quote, urlparse import aiohttp from aiocache import cached @@ -772,6 +772,21 @@ def is_openai_new_model(model: str) -> bool: return False +def _sanitize_model_for_url(model: str) -> str: + """Sanitize a model name before interpolating it into a URL path. + + Rejects path traversal attempts (../, /, \\) and percent-encodes + the name so it is safe to use as a single URL path segment + (e.g. Azure deployment name). + """ + if not model or '..' in model or '/' in model or '\\' in model: + raise HTTPException( + status_code=400, + detail='Invalid model name: must not be empty or contain path separators or traversal sequences', + ) + return quote(model, safe='') + + def convert_to_azure_payload(url, payload: dict, api_version: str): model = payload.get('model', '') @@ -795,6 +810,9 @@ def convert_to_azure_payload(url, payload: dict, api_version: str): # Filter out unsupported parameters payload = {k: v for k, v in payload.items() if k in allowed_params} + # Sanitize model name to prevent path traversal in the deployment URL + model = _sanitize_model_for_url(model) + url = f'{url}/openai/deployments/{model}' return url, payload @@ -1364,7 +1382,7 @@ async def responses( else: api_version = api_config.get('api_version', '2023-03-15-preview') headers['api-version'] = api_version - model = payload.get('model', '') + model = _sanitize_model_for_url(payload.get('model', '')) request_url = f'{url}/openai/deployments/{model}/responses?api-version={api_version}' else: request_url = f'{url}/responses' @@ -1404,6 +1422,8 @@ async def responses( return response_data + except HTTPException: + raise except Exception as e: log.exception(e) raise HTTPException( @@ -1515,6 +1535,8 @@ async def proxy(path: str, request: Request, user=Depends(get_verified_user)): return response_data + except HTTPException: + raise except Exception as e: log.exception(e) raise HTTPException( From 15f9a8f3f13f112c96cb1b16f88859f65de58346 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Sun, 12 Apr 2026 12:36:21 -0500 Subject: [PATCH 29/67] refac --- backend/open_webui/routers/audio.py | 14 +++++++------- backend/open_webui/routers/ollama.py | 2 +- backend/open_webui/storage/provider.py | 12 ++++++------ 3 files changed, 14 insertions(+), 14 deletions(-) diff --git a/backend/open_webui/routers/audio.py b/backend/open_webui/routers/audio.py index 9d8938b419..b1886b2137 100644 --- a/backend/open_webui/routers/audio.py +++ b/backend/open_webui/routers/audio.py @@ -660,7 +660,7 @@ def transcription_handler(request, file_path, metadata, user=None): data = {'text': transcript.strip()} # save the transcript to a json file - transcript_file = f'{file_dir}/{id}.json' + transcript_file = os.path.join(file_dir, f'{id}.json') with open(transcript_file, 'w') as f: json.dump(data, f) @@ -698,7 +698,7 @@ def transcription_handler(request, file_path, metadata, user=None): data = r.json() # save the transcript to a json file - transcript_file = f'{file_dir}/{id}.json' + transcript_file = os.path.join(file_dir, f'{id}.json') with open(transcript_file, 'w') as f: json.dump(data, f) @@ -767,7 +767,7 @@ def transcription_handler(request, file_path, metadata, user=None): data = {'text': transcript.strip()} # Save transcript - transcript_file = f'{file_dir}/{id}.json' + transcript_file = os.path.join(file_dir, f'{id}.json') with open(transcript_file, 'w') as f: json.dump(data, f) @@ -874,7 +874,7 @@ def transcription_handler(request, file_path, metadata, user=None): data = {'text': transcript} # Save transcript to json file (consistent with other providers) - transcript_file = f'{file_dir}/{id}.json' + transcript_file = os.path.join(file_dir, f'{id}.json') with open(transcript_file, 'w') as f: json.dump(data, f) @@ -1059,7 +1059,7 @@ def transcription_handler(request, file_path, metadata, user=None): data = {'text': transcript} # Save transcript to json file (consistent with other providers) - transcript_file = f'{file_dir}/{id}.json' + transcript_file = os.path.join(file_dir, f'{id}.json') with open(transcript_file, 'w') as f: json.dump(data, f) @@ -1237,9 +1237,9 @@ def transcription( filename = f'{id}.{ext}' contents = file.file.read() - file_dir = f'{CACHE_DIR}/audio/transcriptions' + file_dir = os.path.join(CACHE_DIR, 'audio', 'transcriptions') os.makedirs(file_dir, exist_ok=True) - file_path = f'{file_dir}/{filename}' + file_path = os.path.join(file_dir, filename) # Defense-in-depth: ensure resolved path stays within intended directory if not os.path.realpath(file_path).startswith(os.path.realpath(file_dir)): diff --git a/backend/open_webui/routers/ollama.py b/backend/open_webui/routers/ollama.py index 7d5a316fbd..5cf87854ac 100644 --- a/backend/open_webui/routers/ollama.py +++ b/backend/open_webui/routers/ollama.py @@ -1585,7 +1585,7 @@ async def download_model( file_name = parse_huggingface_url(form_data.url) if file_name: - file_path = f'{UPLOAD_DIR}/{file_name}' + file_path = os.path.join(UPLOAD_DIR, file_name) return StreamingResponse( download_file_stream(url, form_data.url, file_path, file_name), diff --git a/backend/open_webui/storage/provider.py b/backend/open_webui/storage/provider.py index 3c29462349..0886c851f0 100644 --- a/backend/open_webui/storage/provider.py +++ b/backend/open_webui/storage/provider.py @@ -61,7 +61,7 @@ class LocalStorageProvider(StorageProvider): contents = file.read() if not contents: raise ValueError(ERROR_MESSAGES.EMPTY_CONTENT) - file_path = f'{UPLOAD_DIR}/{filename}' + file_path = os.path.join(UPLOAD_DIR, filename) with open(file_path, 'wb') as f: f.write(contents) return contents, file_path @@ -74,8 +74,8 @@ class LocalStorageProvider(StorageProvider): @staticmethod def delete_file(file_path: str) -> None: """Handles deletion of the file from local storage.""" - filename = file_path.split('/')[-1] - file_path = f'{UPLOAD_DIR}/{filename}' + filename = os.path.basename(file_path) + file_path = os.path.join(UPLOAD_DIR, filename) if os.path.isfile(file_path): os.remove(file_path) else: @@ -202,7 +202,7 @@ class S3StorageProvider(StorageProvider): return '/'.join(full_file_path.split('//')[1].split('/')[1:]) def _get_local_file_path(self, s3_key: str) -> str: - return f'{UPLOAD_DIR}/{s3_key.split("/")[-1]}' + return os.path.join(UPLOAD_DIR, s3_key.split('/')[-1]) class GCSStorageProvider(StorageProvider): @@ -234,7 +234,7 @@ class GCSStorageProvider(StorageProvider): """Handles downloading of the file from GCS storage.""" try: filename = file_path.removeprefix('gs://').split('/')[1] - local_file_path = f'{UPLOAD_DIR}/{filename}' + local_file_path = os.path.join(UPLOAD_DIR, filename) blob = self.bucket.get_blob(filename) blob.download_to_filename(local_file_path) @@ -298,7 +298,7 @@ class AzureStorageProvider(StorageProvider): """Handles downloading of the file from Azure Blob Storage.""" try: filename = file_path.split('/')[-1] - local_file_path = f'{UPLOAD_DIR}/{filename}' + local_file_path = os.path.join(UPLOAD_DIR, filename) blob_client = self.container_client.get_blob_client(filename) with open(local_file_path, 'wb') as download_file: download_file.write(blob_client.download_blob().readall()) From b618d840657fe15377ef0fe5a40223825e9d8b16 Mon Sep 17 00:00:00 2001 From: Classic298 <27028174+Classic298@users.noreply.github.com> Date: Sun, 12 Apr 2026 19:41:51 +0200 Subject: [PATCH 30/67] fix: add missing read-access check on channel members endpoint (#23625) The GET /channels/{id}/members endpoint checked membership for group/dm channels but had no access gate for standard channels, allowing any authenticated user with channels permission to enumerate members of private standard channels by UUID. --- backend/open_webui/routers/channels.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/backend/open_webui/routers/channels.py b/backend/open_webui/routers/channels.py index 714bde8a85..71bb34394f 100644 --- a/backend/open_webui/routers/channels.py +++ b/backend/open_webui/routers/channels.py @@ -467,6 +467,9 @@ async def get_channel_members_by_id( if channel.type in ['group', 'dm']: if not Channels.is_user_channel_member(channel.id, user.id, db=db): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) + else: + if user.role != 'admin' and not channel_has_access(user.id, channel, permission='read', db=db): + raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) if channel.type == 'dm': user_ids = [member.user_id for member in Channels.get_members_by_channel_id(channel.id, db=db)] From 27169124f220e5cea21c88601c731c3749496ab0 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Sun, 12 Apr 2026 14:22:11 -0500 Subject: [PATCH 31/67] refac: async db --- backend/open_webui/functions.py | 24 +- backend/open_webui/internal/db.py | 111 ++- backend/open_webui/main.py | 49 +- backend/open_webui/models/access_grants.py | 190 ++-- backend/open_webui/models/auths.py | 78 +- backend/open_webui/models/automations.py | 160 ++-- backend/open_webui/models/channels.py | 486 +++++----- backend/open_webui/models/chat_messages.py | 331 +++---- backend/open_webui/models/chats.py | 853 +++++++++--------- backend/open_webui/models/feedbacks.py | 232 ++--- backend/open_webui/models/files.py | 188 ++-- backend/open_webui/models/folders.py | 163 ++-- backend/open_webui/models/functions.py | 203 ++--- backend/open_webui/models/groups.py | 307 ++++--- backend/open_webui/models/knowledge.py | 343 +++---- backend/open_webui/models/memories.py | 75 +- backend/open_webui/models/messages.py | 290 +++--- backend/open_webui/models/models.py | 269 +++--- backend/open_webui/models/notes.py | 153 ++-- backend/open_webui/models/oauth_sessions.py | 177 ++-- backend/open_webui/models/prompt_history.py | 100 +- backend/open_webui/models/prompts.py | 284 +++--- backend/open_webui/models/skills.py | 168 ++-- backend/open_webui/models/tags.py | 66 +- backend/open_webui/models/tools.py | 138 +-- backend/open_webui/models/users.py | 405 +++++---- backend/open_webui/retrieval/utils.py | 18 +- backend/open_webui/routers/analytics.py | 58 +- backend/open_webui/routers/audio.py | 6 +- backend/open_webui/routers/auths.py | 132 +-- backend/open_webui/routers/automations.py | 90 +- backend/open_webui/routers/channels.py | 450 ++++----- backend/open_webui/routers/chats.py | 264 +++--- backend/open_webui/routers/evaluations.py | 66 +- backend/open_webui/routers/files.py | 106 +-- backend/open_webui/routers/folders.py | 70 +- backend/open_webui/routers/functions.py | 100 +- backend/open_webui/routers/groups.py | 62 +- backend/open_webui/routers/images.py | 32 +- backend/open_webui/routers/knowledge.py | 170 ++-- backend/open_webui/routers/memories.py | 46 +- backend/open_webui/routers/models.py | 112 +-- backend/open_webui/routers/notes.py | 74 +- backend/open_webui/routers/ollama.py | 62 +- backend/open_webui/routers/openai.py | 20 +- backend/open_webui/routers/prompts.py | 142 +-- backend/open_webui/routers/retrieval.py | 56 +- backend/open_webui/routers/scim.py | 140 +-- backend/open_webui/routers/skills.py | 82 +- backend/open_webui/routers/terminals.py | 14 +- backend/open_webui/routers/tools.py | 138 +-- backend/open_webui/routers/users.py | 110 +-- backend/open_webui/socket/main.py | 70 +- backend/open_webui/tools/builtin.py | 164 ++-- .../utils/access_control/__init__.py | 42 +- .../open_webui/utils/access_control/files.py | 22 +- backend/open_webui/utils/actions.py | 14 +- backend/open_webui/utils/auth.py | 26 +- backend/open_webui/utils/automations.py | 18 +- backend/open_webui/utils/chat.py | 12 +- backend/open_webui/utils/embeddings.py | 2 +- backend/open_webui/utils/files.py | 55 +- backend/open_webui/utils/filter.py | 39 +- backend/open_webui/utils/groups.py | 4 +- backend/open_webui/utils/mcp/client.py | 2 +- backend/open_webui/utils/middleware.py | 137 +-- backend/open_webui/utils/models.py | 32 +- backend/open_webui/utils/oauth.py | 88 +- backend/open_webui/utils/plugin.py | 40 +- backend/open_webui/utils/redis.py | 2 +- backend/open_webui/utils/telemetry/metrics.py | 12 +- backend/open_webui/utils/tools.py | 62 +- backend/requirements-min.txt | 2 + backend/requirements.txt | 2 + 74 files changed, 4831 insertions(+), 4479 deletions(-) diff --git a/backend/open_webui/functions.py b/backend/open_webui/functions.py index 6a1fe22149..37a9011aab 100644 --- a/backend/open_webui/functions.py +++ b/backend/open_webui/functions.py @@ -53,12 +53,12 @@ logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL) log = logging.getLogger(__name__) -def get_function_module_by_id(request: Request, pipe_id: str): - function_module, _, _ = get_function_module_from_cache(request, pipe_id) +async def get_function_module_by_id(request: Request, pipe_id: str): + function_module, _, _ = await get_function_module_from_cache(request, pipe_id) if hasattr(function_module, 'valves') and hasattr(function_module, 'Valves'): Valves = function_module.Valves - valves = Functions.get_function_valves_by_id(pipe_id) + valves = await Functions.get_function_valves_by_id(pipe_id) if valves: try: @@ -73,12 +73,12 @@ def get_function_module_by_id(request: Request, pipe_id: str): async def get_function_models(request): - pipes = Functions.get_functions_by_type('pipe', active_only=True) + pipes = await Functions.get_functions_by_type('pipe', active_only=True) pipe_models = [] for pipe in pipes: try: - function_module = get_function_module_by_id(request, pipe.id) + function_module = await get_function_module_by_id(request, pipe.id) has_user_valves = False if hasattr(function_module, 'UserValves'): @@ -187,7 +187,7 @@ async def generate_function_chat_completion(request, form_data, user, models: di pipe_id, _ = pipe_id.split('.', 1) return pipe_id - def get_function_params(function_module, form_data, user, extra_params=None): + async def get_function_params(function_module, form_data, user, extra_params=None): if extra_params is None: extra_params = {} @@ -198,7 +198,7 @@ async def generate_function_chat_completion(request, form_data, user, models: di params = {'body': form_data} | {k: v for k, v in extra_params.items() if k in sig.parameters} if '__user__' in params and hasattr(function_module, 'UserValves'): - user_valves = Functions.get_user_valves_by_id_and_user_id(pipe_id, user.id) + user_valves = await Functions.get_user_valves_by_id_and_user_id(pipe_id, user.id) try: params['__user__']['valves'] = function_module.UserValves(**user_valves) except Exception as e: @@ -208,7 +208,7 @@ async def generate_function_chat_completion(request, form_data, user, models: di return params model_id = form_data.get('model') - model_info = Models.get_model_by_id(model_id) + model_info = await Models.get_model_by_id(model_id) metadata = form_data.pop('metadata', {}) @@ -225,8 +225,8 @@ async def generate_function_chat_completion(request, form_data, user, models: di if metadata: if all(k in metadata for k in ('session_id', 'chat_id', 'message_id')): - __event_emitter__ = get_event_emitter(metadata) - __event_call__ = get_event_call(metadata) + __event_emitter__ = await get_event_emitter(metadata) + __event_call__ = await get_event_call(metadata) __task__ = metadata.get('task', None) __task_body__ = metadata.get('task_body', None) @@ -268,10 +268,10 @@ async def generate_function_chat_completion(request, form_data, user, models: di form_data = apply_system_prompt_to_body(system, form_data, metadata, user) pipe_id = get_pipe_id(form_data) - function_module = get_function_module_by_id(request, pipe_id) + function_module = await get_function_module_by_id(request, pipe_id) pipe = function_module.pipe - params = get_function_params(function_module, form_data, user, extra_params) + params = await get_function_params(function_module, form_data, user, extra_params) if form_data.get('stream', False): diff --git a/backend/open_webui/internal/db.py b/backend/open_webui/internal/db.py index b0545255a6..a9e5e089ab 100644 --- a/backend/open_webui/internal/db.py +++ b/backend/open_webui/internal/db.py @@ -1,7 +1,7 @@ import os import json import logging -from contextlib import contextmanager +from contextlib import asynccontextmanager, contextmanager from typing import Any, Optional from open_webui.internal.wrappers import register_connection @@ -19,6 +19,7 @@ from open_webui.env import ( ) from peewee_migrate import Router from sqlalchemy import Dialect, create_engine, MetaData, event, types +from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker from sqlalchemy.ext.declarative import declarative_base from sqlalchemy.orm import scoped_session, sessionmaker, Session from sqlalchemy.pool import QueuePool, NullPool @@ -81,6 +82,32 @@ if ENABLE_DB_MIGRATIONS: SQLALCHEMY_DATABASE_URL = DATABASE_URL + +def _make_async_url(url: str) -> str: + """Convert a sync database URL to its async driver equivalent.""" + if url.startswith('sqlite+sqlcipher://'): + # SQLCipher has no async driver — not supported for async + raise ValueError( + 'sqlite+sqlcipher:// URLs are not supported with async engine. ' + 'Use standard sqlite:// or postgresql:// instead.' + ) + if url.startswith('sqlite:///') or url.startswith('sqlite://'): + return url.replace('sqlite://', 'sqlite+aiosqlite://', 1) + if url.startswith('postgresql+psycopg2://'): + return url.replace('postgresql+psycopg2://', 'postgresql+asyncpg://', 1) + if url.startswith('postgresql://'): + return url.replace('postgresql://', 'postgresql+asyncpg://', 1) + if url.startswith('postgres://'): + return url.replace('postgres://', 'postgresql+asyncpg://', 1) + # For other dialects, return as-is and let SQLAlchemy handle it + return url + + +# ============================================================ +# SYNC ENGINE (used only for: startup migrations, config loading, +# Alembic, peewee migration, health checks) +# ============================================================ + # Handle SQLCipher URLs if SQLALCHEMY_DATABASE_URL.startswith('sqlite+sqlcipher://'): database_password = os.environ.get('DATABASE_PASSWORD') @@ -155,6 +182,7 @@ else: engine = create_engine(SQLALCHEMY_DATABASE_URL, pool_pre_ping=True) +# Sync session — used ONLY for startup config loading (config.py runs at import time) SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine, expire_on_commit=False) metadata_obj = MetaData(schema=DATABASE_SCHEMA) Base = declarative_base(metadata=metadata_obj) @@ -162,6 +190,7 @@ ScopedSession = scoped_session(SessionLocal) def get_session(): + """Sync session generator — used ONLY for startup/config operations.""" db = SessionLocal() try: yield db @@ -172,10 +201,82 @@ def get_session(): get_db = contextmanager(get_session) -@contextmanager -def get_db_context(db: Optional[Session] = None): - if isinstance(db, Session) and DATABASE_ENABLE_SESSION_SHARING: +# ============================================================ +# ASYNC ENGINE (used for ALL runtime database operations) +# ============================================================ + +ASYNC_SQLALCHEMY_DATABASE_URL = _make_async_url(SQLALCHEMY_DATABASE_URL) + +if 'sqlite' in ASYNC_SQLALCHEMY_DATABASE_URL: + async_engine = create_async_engine( + ASYNC_SQLALCHEMY_DATABASE_URL, + connect_args={'check_same_thread': False}, + ) + + if DATABASE_ENABLE_SQLITE_WAL: + @event.listens_for(async_engine.sync_engine, 'connect') + def _set_sqlite_wal(dbapi_connection, connection_record): + cursor = dbapi_connection.cursor() + cursor.execute('PRAGMA journal_mode=WAL') + cursor.close() +else: + if isinstance(DATABASE_POOL_SIZE, int): + if DATABASE_POOL_SIZE > 0: + async_engine = create_async_engine( + ASYNC_SQLALCHEMY_DATABASE_URL, + pool_size=DATABASE_POOL_SIZE, + max_overflow=DATABASE_POOL_MAX_OVERFLOW, + pool_timeout=DATABASE_POOL_TIMEOUT, + pool_recycle=DATABASE_POOL_RECYCLE, + pool_pre_ping=True, + poolclass=QueuePool, + ) + else: + async_engine = create_async_engine( + ASYNC_SQLALCHEMY_DATABASE_URL, + pool_pre_ping=True, + poolclass=NullPool, + ) + else: + async_engine = create_async_engine( + ASYNC_SQLALCHEMY_DATABASE_URL, + pool_pre_ping=True, + ) + + +AsyncSessionLocal = async_sessionmaker( + bind=async_engine, + class_=AsyncSession, + autocommit=False, + autoflush=False, + expire_on_commit=False, +) + + +async def get_async_session(): + """Async session generator for FastAPI Depends().""" + async with AsyncSessionLocal() as db: + try: + yield db + finally: + await db.close() + + +@asynccontextmanager +async def get_async_db(): + """Async context manager for use outside of FastAPI dependency injection.""" + async with AsyncSessionLocal() as db: + try: + yield db + finally: + await db.close() + + +@asynccontextmanager +async def get_async_db_context(db: Optional[AsyncSession] = None): + """Async context manager that reuses an existing session if provided and session sharing is enabled.""" + if isinstance(db, AsyncSession) and DATABASE_ENABLE_SESSION_SHARING: yield db else: - with get_db() as session: + async with get_async_db() as session: yield session diff --git a/backend/open_webui/main.py b/backend/open_webui/main.py index 8fc56a6659..d959351d4c 100644 --- a/backend/open_webui/main.py +++ b/backend/open_webui/main.py @@ -108,8 +108,8 @@ from open_webui.routers.retrieval import ( ) -from sqlalchemy.orm import Session -from open_webui.internal.db import ScopedSession, engine, get_session +from sqlalchemy.ext.asyncio import AsyncSession +from open_webui.internal.db import ScopedSession, engine, get_async_session from open_webui.models.functions import Functions from open_webui.models.models import Models @@ -575,7 +575,7 @@ from open_webui.constants import ERROR_MESSAGES if SAFE_MODE: print('SAFE MODE ENABLED') - Functions.deactivate_all_functions() + # Functions.deactivate_all_functions() is awaited in lifespan below logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL) log = logging.getLogger(__name__) @@ -629,14 +629,17 @@ async def lifespan(app: FastAPI): # Create admin account from env vars if specified and no users exist if WEBUI_ADMIN_EMAIL and WEBUI_ADMIN_PASSWORD: - if create_admin_user(WEBUI_ADMIN_EMAIL, WEBUI_ADMIN_PASSWORD, WEBUI_ADMIN_NAME): + if await create_admin_user(WEBUI_ADMIN_EMAIL, WEBUI_ADMIN_PASSWORD, WEBUI_ADMIN_NAME): # Disable signup since we now have an admin app.state.config.ENABLE_SIGNUP = False + if SAFE_MODE: + await Functions.deactivate_all_functions() + # This should be blocking (sync) so functions are not deactivated on first /get_models calls # when the first user lands on the / route. log.info('Installing external dependencies of functions and tools...') - install_tool_and_function_dependencies() + await install_tool_and_function_dependencies() app.state.redis = get_redis_connection( redis_url=REDIS_URL, @@ -1605,7 +1608,7 @@ async def get_models(request: Request, refresh: bool = False, user=Depends(get_v ) ) - models = get_filtered_models(models, user) + models = await get_filtered_models(models, user) log.debug( f'/api/models returned filtered models accessible to the user: {json.dumps([model.get("id") for model in models])}' @@ -1671,12 +1674,12 @@ async def chat_completion( raise Exception('Model not found') model = request.app.state.MODELS[model_id] - model_info = Models.get_model_by_id(model_id) + model_info = await Models.get_model_by_id(model_id) # Check if user has access to the model if not BYPASS_MODEL_ACCESS_CONTROL and (user.role != 'admin' or not BYPASS_ADMIN_ACCESS_CONTROL): try: - check_model_access(user, model) + await check_model_access(user, model) except Exception as e: raise e else: @@ -1758,7 +1761,7 @@ async def chat_completion( # Verify chat ownership — lightweight EXISTS check avoids # deserializing the full chat JSON blob just to confirm the row exists if ( - not Chats.is_chat_owner(metadata['chat_id'], user.id) and user.role != 'admin' + not await Chats.is_chat_owner(metadata['chat_id'], user.id) and user.role != 'admin' ): # admins can access any chat raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -1770,7 +1773,7 @@ async def chat_completion( parent_message_files = parent_message.get('files', []) if parent_message_files: try: - Chats.insert_chat_files( + await Chats.insert_chat_files( metadata['chat_id'], parent_message.get('id'), [ @@ -1802,7 +1805,7 @@ async def chat_completion( if metadata.get('chat_id') and metadata.get('message_id'): try: if not metadata['chat_id'].startswith('local:'): - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], { @@ -1813,13 +1816,13 @@ async def chat_completion( except Exception: pass - ctx = build_chat_response_context(request, form_data, user, model, metadata, tasks, events) + ctx = await build_chat_response_context(request, form_data, user, model, metadata, tasks, events) return await process_chat_response(response, ctx) except asyncio.CancelledError: log.info('Chat processing was cancelled') try: - event_emitter = get_event_emitter(metadata) + event_emitter = await get_event_emitter(metadata) await asyncio.shield( event_emitter( {'type': 'chat:tasks:cancel'}, @@ -1835,7 +1838,7 @@ async def chat_completion( # Update the chat message with the error try: if not metadata['chat_id'].startswith('local:'): - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], { @@ -1844,7 +1847,7 @@ async def chat_completion( }, ) - event_emitter = get_event_emitter(metadata) + event_emitter = await get_event_emitter(metadata) await event_emitter( { 'type': 'chat:message:error', @@ -1883,7 +1886,7 @@ async def chat_completion( # Emit chat:active=false when task completes try: if metadata.get('chat_id'): - event_emitter = get_event_emitter(metadata, update_db=False) + event_emitter = await get_event_emitter(metadata, update_db=False) if event_emitter: await event_emitter({'type': 'chat:active', 'data': {'active': False}}) except Exception as e: @@ -1897,7 +1900,7 @@ async def chat_completion( id=metadata['chat_id'], ) # Emit chat:active=true when task starts - event_emitter = get_event_emitter(metadata, update_db=False) + event_emitter = await get_event_emitter(metadata, update_db=False) if event_emitter: await event_emitter({'type': 'chat:active', 'data': {'active': True}}) return {'status': True, 'task_id': task_id} @@ -2024,7 +2027,7 @@ async def list_tasks_endpoint(request: Request, user=Depends(get_verified_user)) @app.get('/api/tasks/chat/{chat_id}') async def list_tasks_by_chat_id_endpoint(request: Request, chat_id: str, user=Depends(get_verified_user)): - chat = Chats.get_chat_by_id(chat_id) + chat = await Chats.get_chat_by_id(chat_id) if chat is None or chat.user_id != user.id: return {'task_ids': []} @@ -2065,9 +2068,9 @@ async def get_app_config(request: Request): detail='Invalid token', ) if data is not None and 'id' in data: - user = Users.get_user_by_id(data['id']) + user = await Users.get_user_by_id(data['id']) - user_count = Users.get_num_users() + user_count = await Users.get_num_users() onboarding = False if user is None: @@ -2276,7 +2279,7 @@ async def get_current_usage(user=Depends(get_verified_user)): return { 'model_ids': get_models_in_use(), - 'user_count': Users.get_active_user_count(), + 'user_count': await Users.get_active_user_count(), } except HTTPException: raise @@ -2483,7 +2486,7 @@ async def oauth_login_callback( provider: str, request: Request, response: Response, - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): return await oauth_manager.handle_callback(request, provider, response, db=db) @@ -2496,7 +2499,7 @@ async def oauth_login_callback( @app.post('/oauth/backchannel-logout') async def oauth_backchannel_logout( request: Request, - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): if not ENABLE_OAUTH_BACKCHANNEL_LOGOUT: raise HTTPException(status_code=404) diff --git a/backend/open_webui/models/access_grants.py b/backend/open_webui/models/access_grants.py index 20601fd30e..f064306a2c 100644 --- a/backend/open_webui/models/access_grants.py +++ b/backend/open_webui/models/access_grants.py @@ -3,8 +3,9 @@ import time import uuid from typing import Optional -from sqlalchemy.orm import Session -from open_webui.internal.db import Base, get_db_context +from sqlalchemy import select, delete +from sqlalchemy.ext.asyncio import AsyncSession +from open_webui.internal.db import Base, get_async_db_context from pydantic import BaseModel, ConfigDict from sqlalchemy import BigInteger, Column, Text, UniqueConstraint, or_, and_ @@ -281,20 +282,20 @@ def grants_to_access_control(grants: list) -> Optional[dict]: class AccessGrantsTable: - def grant_access( + async def grant_access( self, resource_type: str, resource_id: str, principal_type: str, principal_id: str, permission: str, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[AccessGrantModel]: """Add a single access grant. Idempotent (ignores duplicates).""" - with get_db_context(db) as db: + async with get_async_db_context(db) as db: # Check for existing grant - existing = ( - db.query(AccessGrant) + result = await db.execute( + select(AccessGrant) .filter_by( resource_type=resource_type, resource_id=resource_id, @@ -302,8 +303,8 @@ class AccessGrantsTable: principal_id=principal_id, permission=permission, ) - .first() ) + existing = result.scalars().first() if existing: return AccessGrantModel.model_validate(existing) @@ -317,23 +318,23 @@ class AccessGrantsTable: created_at=int(time.time()), ) db.add(grant) - db.commit() - db.refresh(grant) + await db.commit() + await db.refresh(grant) return AccessGrantModel.model_validate(grant) - def revoke_access( + async def revoke_access( self, resource_type: str, resource_id: str, principal_type: str, principal_id: str, permission: str, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> bool: """Remove a single access grant.""" - with get_db_context(db) as db: - deleted = ( - db.query(AccessGrant) + async with get_async_db_context(db) as db: + result = await db.execute( + delete(AccessGrant) .filter_by( resource_type=resource_type, resource_id=resource_id, @@ -341,47 +342,47 @@ class AccessGrantsTable: principal_id=principal_id, permission=permission, ) - .delete() ) - db.commit() - return deleted > 0 + await db.commit() + return result.rowcount > 0 - def revoke_all_access( + async def revoke_all_access( self, resource_type: str, resource_id: str, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> int: """Remove all access grants for a resource.""" - with get_db_context(db) as db: - deleted = ( - db.query(AccessGrant) + async with get_async_db_context(db) as db: + result = await db.execute( + delete(AccessGrant) .filter_by( resource_type=resource_type, resource_id=resource_id, ) - .delete() ) - db.commit() - return deleted + await db.commit() + return result.rowcount - def set_access_control( + async def set_access_control( self, resource_type: str, resource_id: str, access_control: Optional[dict], - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[AccessGrantModel]: """ Replace all grants for a resource from an access_control JSON dict. This is the primary bridge for backward compat with the frontend. """ - with get_db_context(db) as db: + async with get_async_db_context(db) as db: # Delete all existing grants for this resource - db.query(AccessGrant).filter_by( - resource_type=resource_type, - resource_id=resource_id, - ).delete() + await db.execute( + delete(AccessGrant).filter_by( + resource_type=resource_type, + resource_id=resource_id, + ) + ) # Convert JSON to grant dicts grant_dicts = access_control_to_grants(resource_type, resource_id, access_control) @@ -397,25 +398,27 @@ class AccessGrantsTable: db.add(grant) results.append(grant) - db.commit() + await db.commit() return [AccessGrantModel.model_validate(g) for g in results] - def set_access_grants( + async def set_access_grants( self, resource_type: str, resource_id: str, access_grants: Optional[list], - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[AccessGrantModel]: """ Replace all grants for a resource from a direct access_grants list. """ - with get_db_context(db) as db: - db.query(AccessGrant).filter_by( - resource_type=resource_type, - resource_id=resource_id, - ).delete() + async with get_async_db_context(db) as db: + await db.execute( + delete(AccessGrant).filter_by( + resource_type=resource_type, + resource_id=resource_id, + ) + ) normalized_grants = normalize_access_grants(access_grants) @@ -433,80 +436,80 @@ class AccessGrantsTable: db.add(grant) results.append(grant) - db.commit() + await db.commit() return [AccessGrantModel.model_validate(g) for g in results] - def get_access_control( + async def get_access_control( self, resource_type: str, resource_id: str, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[dict]: """ Reconstruct the old-style access_control JSON dict from grants. For backward compat with the frontend. """ - with get_db_context(db) as db: - grants = ( - db.query(AccessGrant) + async with get_async_db_context(db) as db: + result = await db.execute( + select(AccessGrant) .filter_by( resource_type=resource_type, resource_id=resource_id, ) - .all() ) + grants = result.scalars().all() grant_models = [AccessGrantModel.model_validate(g) for g in grants] return grants_to_access_control(grant_models) - def get_grants_by_resource( + async def get_grants_by_resource( self, resource_type: str, resource_id: str, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[AccessGrantModel]: """Get all grants for a specific resource.""" - with get_db_context(db) as db: - grants = ( - db.query(AccessGrant) + async with get_async_db_context(db) as db: + result = await db.execute( + select(AccessGrant) .filter_by( resource_type=resource_type, resource_id=resource_id, ) - .all() ) + grants = result.scalars().all() return [AccessGrantModel.model_validate(g) for g in grants] - def get_grants_by_resources( + async def get_grants_by_resources( self, resource_type: str, resource_ids: list[str], - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> dict[str, list[AccessGrantModel]]: """Batch-fetch grants for multiple resources. Returns {resource_id: [grants]}.""" if not resource_ids: return {} - with get_db_context(db) as db: - grants = ( - db.query(AccessGrant) + async with get_async_db_context(db) as db: + result = await db.execute( + select(AccessGrant) .filter( AccessGrant.resource_type == resource_type, AccessGrant.resource_id.in_(resource_ids), ) - .all() ) - result: dict[str, list[AccessGrantModel]] = {rid: [] for rid in resource_ids} + grants = result.scalars().all() + result_dict: dict[str, list[AccessGrantModel]] = {rid: [] for rid in resource_ids} for g in grants: - result[g.resource_id].append(AccessGrantModel.model_validate(g)) - return result + result_dict[g.resource_id].append(AccessGrantModel.model_validate(g)) + return result_dict - def has_access( + async def has_access( self, user_id: str, resource_type: str, resource_id: str, permission: str = 'read', user_group_ids: Optional[set[str]] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> bool: """ Check if a user has the specified permission on a resource. @@ -516,7 +519,7 @@ class AccessGrantsTable: - There's a grant for the specific user with the requested permission - There's a grant for any of the user's groups with the requested permission """ - with get_db_context(db) as db: + async with get_async_db_context(db) as db: # Build conditions for matching grants conditions = [ # Public access @@ -535,7 +538,7 @@ class AccessGrantsTable: if user_group_ids is None: from open_webui.models.groups import Groups - user_groups = Groups.get_groups_by_member_id(user_id, db=db) + user_groups = await Groups.get_groups_by_member_id(user_id, db=db) user_group_ids = {group.id for group in user_groups} if user_group_ids: @@ -546,26 +549,27 @@ class AccessGrantsTable: ) ) - exists = ( - db.query(AccessGrant) + result = await db.execute( + select(AccessGrant) .filter( AccessGrant.resource_type == resource_type, AccessGrant.resource_id == resource_id, AccessGrant.permission == permission, or_(*conditions), ) - .first() + .limit(1) ) - return exists is not None + grant = result.scalars().first() + return grant is not None - def get_accessible_resource_ids( + async def get_accessible_resource_ids( self, user_id: str, resource_type: str, resource_ids: list[str], permission: str = 'read', user_group_ids: Optional[set[str]] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> set[str]: """ Batch check: return the subset of resource_ids that the user can access. @@ -575,7 +579,7 @@ class AccessGrantsTable: if not resource_ids: return set() - with get_db_context(db) as db: + async with get_async_db_context(db) as db: conditions = [ and_( AccessGrant.principal_type == 'user', @@ -590,7 +594,7 @@ class AccessGrantsTable: if user_group_ids is None: from open_webui.models.groups import Groups - user_groups = Groups.get_groups_by_member_id(user_id, db=db) + user_groups = await Groups.get_groups_by_member_id(user_id, db=db) user_group_ids = {group.id for group in user_groups} if user_group_ids: @@ -601,8 +605,8 @@ class AccessGrantsTable: ) ) - rows = ( - db.query(AccessGrant.resource_id) + result = await db.execute( + select(AccessGrant.resource_id) .filter( AccessGrant.resource_type == resource_type, AccessGrant.resource_id.in_(resource_ids), @@ -610,16 +614,16 @@ class AccessGrantsTable: or_(*conditions), ) .distinct() - .all() ) + rows = result.all() return {row[0] for row in rows} - def get_users_with_access( + async def get_users_with_access( self, resource_type: str, resource_id: str, permission: str = 'read', - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list: """ Get all users who have the specified permission on a resource. @@ -628,21 +632,21 @@ class AccessGrantsTable: from open_webui.models.users import Users, UserModel from open_webui.models.groups import Groups - with get_db_context(db) as db: - grants = ( - db.query(AccessGrant) + async with get_async_db_context(db) as db: + result = await db.execute( + select(AccessGrant) .filter_by( resource_type=resource_type, resource_id=resource_id, permission=permission, ) - .all() ) + grants = result.scalars().all() # Check for public access for grant in grants: if grant.principal_type == 'user' and grant.principal_id == '*': - result = Users.get_users(filter={'roles': ['!pending']}, db=db) + result = await Users.get_users(filter={'roles': ['!pending']}, db=db) return result.get('users', []) user_ids_with_access = set() @@ -651,14 +655,14 @@ class AccessGrantsTable: if grant.principal_type == 'user': user_ids_with_access.add(grant.principal_id) elif grant.principal_type == 'group': - group_user_ids = Groups.get_group_user_ids_by_id(grant.principal_id, db=db) + group_user_ids = await Groups.get_group_user_ids_by_id(grant.principal_id, db=db) if group_user_ids: user_ids_with_access.update(group_user_ids) if not user_ids_with_access: return [] - return Users.get_users_by_user_ids(list(user_ids_with_access), db=db) + return await Users.get_users_by_user_ids(list(user_ids_with_access), db=db) def has_permission_filter( self, @@ -673,6 +677,10 @@ class AccessGrantsTable: Apply access control filtering to a SQLAlchemy query by JOINing with access_grant. This replaces the old JSON-column-based filtering with a proper relational JOIN. + + Note: This method builds SQLAlchemy expressions and does NOT perform I/O itself, + so it remains synchronous. The caller is responsible for executing the query + asynchronously with `await db.execute(...)`. """ group_ids = filter.get('group_ids', []) user_id = filter.get('user_id') @@ -718,7 +726,7 @@ class AccessGrantsTable: # LEFT JOIN access_grant and filter # We use a subquery approach to avoid duplicates from multiple matching grants - from sqlalchemy import exists as sa_exists, select + from sqlalchemy import exists as sa_exists grant_exists = ( select(AccessGrant.id) @@ -776,11 +784,15 @@ class AccessGrantsTable: """ Filter for items where user has read BUT NOT write access. Public items are NOT considered read_only. + + Note: This method builds SQLAlchemy expressions and does NOT perform I/O itself, + so it remains synchronous. The caller is responsible for executing the query + asynchronously with `await db.execute(...)`. """ group_ids = filter.get('group_ids', []) user_id = filter.get('user_id') - from sqlalchemy import exists as sa_exists, select + from sqlalchemy import exists as sa_exists # Has read grant (not public) read_grant_exists = ( diff --git a/backend/open_webui/models/auths.py b/backend/open_webui/models/auths.py index 1a1b164c12..ca5070878e 100644 --- a/backend/open_webui/models/auths.py +++ b/backend/open_webui/models/auths.py @@ -2,8 +2,9 @@ import logging import uuid from typing import Optional -from sqlalchemy.orm import Session -from open_webui.internal.db import Base, JSONField, get_db, get_db_context +from sqlalchemy import select, delete, update +from sqlalchemy.ext.asyncio import AsyncSession +from open_webui.internal.db import Base, JSONField, get_async_db_context from open_webui.models.users import User, UserModel, UserProfileImageResponse, Users from open_webui.utils.validate import validate_profile_image_url from pydantic import BaseModel, field_validator @@ -88,7 +89,7 @@ class AddUserForm(SignupForm): class AuthsTable: - def insert_new_auth( + async def insert_new_auth( self, email: str, password: str, @@ -96,9 +97,9 @@ class AuthsTable: profile_image_url: str = '/user.png', role: str = 'pending', oauth: Optional[dict] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[UserModel]: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: log.info('insert_new_auth') id = str(uuid.uuid4()) @@ -107,28 +108,29 @@ class AuthsTable: result = Auth(**auth.model_dump()) db.add(result) - user = Users.insert_new_user(id, name, email, profile_image_url, role, oauth=oauth, db=db) + user = await Users.insert_new_user(id, name, email, profile_image_url, role, oauth=oauth, db=db) - db.commit() - db.refresh(result) + await db.commit() + await db.refresh(result) if result and user: return user else: return None - def authenticate_user( - self, email: str, verify_password: callable, db: Optional[Session] = None + async def authenticate_user( + self, email: str, verify_password: callable, db: Optional[AsyncSession] = None ) -> Optional[UserModel]: log.info(f'authenticate_user: {email}') - user = Users.get_user_by_email(email, db=db) + user = await Users.get_user_by_email(email, db=db) if not user: return None try: - with get_db_context(db) as db: - auth = db.query(Auth).filter_by(id=user.id, active=True).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(Auth).filter_by(id=user.id, active=True)) + auth = result.scalars().first() if auth: if verify_password(auth.password): return user @@ -139,66 +141,66 @@ class AuthsTable: except Exception: return None - def authenticate_user_by_api_key(self, api_key: str, db: Optional[Session] = None) -> Optional[UserModel]: + async def authenticate_user_by_api_key(self, api_key: str, db: Optional[AsyncSession] = None) -> Optional[UserModel]: log.info(f'authenticate_user_by_api_key') # if no api_key, return None if not api_key: return None try: - user = Users.get_user_by_api_key(api_key, db=db) + user = await Users.get_user_by_api_key(api_key, db=db) return user if user else None except Exception: return False - def authenticate_user_by_email(self, email: str, db: Optional[Session] = None) -> Optional[UserModel]: + async def authenticate_user_by_email(self, email: str, db: Optional[AsyncSession] = None) -> Optional[UserModel]: log.info(f'authenticate_user_by_email: {email}') try: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: # Single JOIN query instead of two separate queries - result = ( - db.query(Auth, User) + result = await db.execute( + select(Auth, User) .join(User, Auth.id == User.id) .filter(Auth.email == email, Auth.active == True) - .first() ) - if result: - _, user = result + row = result.first() + if row: + _, user = row return UserModel.model_validate(user) return None except Exception: return None - def update_user_password_by_id(self, id: str, new_password: str, db: Optional[Session] = None) -> bool: + async def update_user_password_by_id(self, id: str, new_password: str, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: - result = db.query(Auth).filter_by(id=id).update({'password': new_password}) - db.commit() - return True if result == 1 else False + async with get_async_db_context(db) as db: + result = await db.execute(update(Auth).filter_by(id=id).values(password=new_password)) + await db.commit() + return True if result.rowcount == 1 else False except Exception: return False - def update_email_by_id(self, id: str, email: str, db: Optional[Session] = None) -> bool: + async def update_email_by_id(self, id: str, email: str, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: - result = db.query(Auth).filter_by(id=id).update({'email': email}) - db.commit() - if result == 1: - Users.update_user_by_id(id, {'email': email}, db=db) + async with get_async_db_context(db) as db: + result = await db.execute(update(Auth).filter_by(id=id).values(email=email)) + await db.commit() + if result.rowcount == 1: + await Users.update_user_by_id(id, {'email': email}, db=db) return True return False except Exception: return False - def delete_auth_by_id(self, id: str, db: Optional[Session] = None) -> bool: + async def delete_auth_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: # Delete User - result = Users.delete_user_by_id(id, db=db) + result = await Users.delete_user_by_id(id, db=db) if result: - db.query(Auth).filter_by(id=id).delete() - db.commit() + await db.execute(delete(Auth).filter_by(id=id)) + await db.commit() return True else: diff --git a/backend/open_webui/models/automations.py b/backend/open_webui/models/automations.py index 7a6bccb9a0..fab3788eb0 100644 --- a/backend/open_webui/models/automations.py +++ b/backend/open_webui/models/automations.py @@ -4,10 +4,10 @@ from typing import Optional from uuid import uuid4 from pydantic import BaseModel, ConfigDict -from sqlalchemy import Column, Text, JSON, Boolean, BigInteger, Index, select, or_, func, cast, String -from sqlalchemy.orm import Session +from sqlalchemy import Column, Text, JSON, Boolean, BigInteger, Index, select, or_, func, cast, String, delete, update +from sqlalchemy.ext.asyncio import AsyncSession -from open_webui.internal.db import Base, get_db, get_db_context +from open_webui.internal.db import Base, get_async_db_context log = logging.getLogger(__name__) @@ -118,14 +118,14 @@ class AutomationListResponse(BaseModel): class AutomationTable: - def insert( + async def insert( self, user_id: str, form: AutomationForm, next_run_at: int, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> AutomationModel: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: now = int(time.time_ns()) row = Automation( id=str(uuid4()), @@ -139,35 +139,38 @@ class AutomationTable: updated_at=now, ) db.add(row) - db.commit() - db.refresh(row) + await db.commit() + await db.refresh(row) return AutomationModel.model_validate(row) - def count_by_user(self, user_id: str, db: Optional[Session] = None) -> int: - with get_db_context(db) as db: - return db.query(Automation).filter_by(user_id=user_id).count() + async def count_by_user(self, user_id: str, db: Optional[AsyncSession] = None) -> int: + async with get_async_db_context(db) as db: + result = await db.execute( + select(func.count()).select_from(Automation).filter_by(user_id=user_id) + ) + return result.scalar() - def get_by_id(self, id: str, db: Optional[Session] = None) -> Optional[AutomationModel]: - with get_db_context(db) as db: - row = db.get(Automation, id) + async def get_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[AutomationModel]: + async with get_async_db_context(db) as db: + row = await db.get(Automation, id) return AutomationModel.model_validate(row) if row else None - def search_automations( + async def search_automations( self, user_id: str, query: Optional[str] = None, status: Optional[str] = None, skip: int = 0, limit: int = 30, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> 'AutomationListResponse': - with get_db_context(db) as db: - q = db.query(Automation).filter_by(user_id=user_id) + async with get_async_db_context(db) as db: + stmt = select(Automation).filter_by(user_id=user_id) if query: search = f'%{query}%' # Search in name and prompt inside JSON data - q = q.filter( + stmt = stmt.filter( or_( Automation.name.ilike(search), cast(Automation.data, String).ilike(search), @@ -175,34 +178,39 @@ class AutomationTable: ) if status == 'active': - q = q.filter(Automation.is_active == True) + stmt = stmt.filter(Automation.is_active == True) elif status == 'paused': - q = q.filter(Automation.is_active == False) + stmt = stmt.filter(Automation.is_active == False) - q = q.order_by(Automation.created_at.desc()) + stmt = stmt.order_by(Automation.created_at.desc()) - total = q.count() + # Get total count + count_result = await db.execute( + select(func.count()).select_from(stmt.subquery()) + ) + total = count_result.scalar() if skip: - q = q.offset(skip) + stmt = stmt.offset(skip) if limit: - q = q.limit(limit) + stmt = stmt.limit(limit) - rows = q.all() + result = await db.execute(stmt) + rows = result.scalars().all() return AutomationListResponse( items=[AutomationModel.model_validate(r) for r in rows], total=total, ) - def update_by_id( + async def update_by_id( self, id: str, form: AutomationForm, next_run_at: int, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[AutomationModel]: - with get_db_context(db) as db: - row = db.get(Automation, id) + async with get_async_db_context(db) as db: + row = await db.get(Automation, id) if not row: return None row.name = form.name @@ -212,37 +220,37 @@ class AutomationTable: row.is_active = form.is_active row.next_run_at = next_run_at row.updated_at = int(time.time_ns()) - db.commit() - db.refresh(row) + await db.commit() + await db.refresh(row) return AutomationModel.model_validate(row) - def toggle( + async def toggle( self, id: str, next_run_at: Optional[int], - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[AutomationModel]: - with get_db_context(db) as db: - row = db.get(Automation, id) + async with get_async_db_context(db) as db: + row = await db.get(Automation, id) if not row: return None row.is_active = not row.is_active row.next_run_at = next_run_at if row.is_active else None row.updated_at = int(time.time_ns()) - db.commit() - db.refresh(row) + await db.commit() + await db.refresh(row) return AutomationModel.model_validate(row) - def delete(self, id: str, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: - row = db.get(Automation, id) + async def delete(self, id: str, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: + row = await db.get(Automation, id) if not row: return False - db.delete(row) - db.commit() + await db.delete(row) + await db.commit() return True - def claim_due(self, now_ns: int, limit: int = 10, db: Optional[Session] = None) -> list[AutomationModel]: + async def claim_due(self, now_ns: int, limit: int = 10, db: Optional[AsyncSession] = None) -> list[AutomationModel]: """ Atomically claim due automations for execution. @@ -250,7 +258,7 @@ class AutomationTable: double-claimed. On PostgreSQL, uses FOR UPDATE SKIP LOCKED for zero-contention distributed work claiming. """ - with get_db_context(db) as db: + async with get_async_db_context(db) as db: stmt = ( select(Automation) .where( @@ -264,7 +272,8 @@ class AutomationTable: if db.bind.dialect.name == 'postgresql': stmt = stmt.with_for_update(skip_locked=True) - rows = db.execute(stmt).scalars().all() + result = await db.execute(stmt) + rows = result.scalars().all() from open_webui.utils.automations import next_run_ns @@ -272,7 +281,7 @@ class AutomationTable: row.last_run_at = now_ns row.next_run_at = next_run_ns(row.data.get('rrule', '')) - db.commit() + await db.commit() return [AutomationModel.model_validate(r) for r in rows] @@ -283,15 +292,15 @@ class AutomationTable: class AutomationRunTable: - def insert( + async def insert( self, automation_id: str, status: str, chat_id: Optional[str] = None, error: Optional[str] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> AutomationRunModel: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: row = AutomationRun( id=str(uuid4()), automation_id=automation_id, @@ -301,30 +310,31 @@ class AutomationRunTable: created_at=int(time.time_ns()), ) db.add(row) - db.commit() - db.refresh(row) + await db.commit() + await db.refresh(row) return AutomationRunModel.model_validate(row) - def get_latest(self, automation_id: str, db: Optional[Session] = None) -> Optional[AutomationRunModel]: - with get_db_context(db) as db: - row = ( - db.query(AutomationRun) + async def get_latest(self, automation_id: str, db: Optional[AsyncSession] = None) -> Optional[AutomationRunModel]: + async with get_async_db_context(db) as db: + result = await db.execute( + select(AutomationRun) .filter_by(automation_id=automation_id) .order_by(AutomationRun.created_at.desc()) - .first() + .limit(1) ) + row = result.scalars().first() return AutomationRunModel.model_validate(row) if row else None - def get_latest_batch( - self, automation_ids: list[str], db: Optional[Session] = None + async def get_latest_batch( + self, automation_ids: list[str], db: Optional[AsyncSession] = None ) -> dict[str, AutomationRunModel]: """Fetch the latest run for each automation in a single query.""" if not automation_ids: return {} - with get_db_context(db) as db: + async with get_async_db_context(db) as db: # Subquery: max created_at per automation_id subq = ( - db.query( + select( AutomationRun.automation_id, func.max(AutomationRun.created_at).label('max_created'), ) @@ -332,43 +342,43 @@ class AutomationRunTable: .group_by(AutomationRun.automation_id) .subquery() ) - rows = ( - db.query(AutomationRun) + result = await db.execute( + select(AutomationRun) .join( subq, (AutomationRun.automation_id == subq.c.automation_id) & (AutomationRun.created_at == subq.c.max_created), ) - .all() ) + rows = result.scalars().all() return { row.automation_id: AutomationRunModel.model_validate(row) for row in rows } - def get_by_automation( + async def get_by_automation( self, automation_id: str, skip: int = 0, limit: int = 50, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[AutomationRunModel]: - with get_db_context(db) as db: - rows = ( - db.query(AutomationRun) + async with get_async_db_context(db) as db: + result = await db.execute( + select(AutomationRun) .filter_by(automation_id=automation_id) .order_by(AutomationRun.created_at.desc()) .offset(skip) .limit(limit) - .all() ) + rows = result.scalars().all() return [AutomationRunModel.model_validate(r) for r in rows] - def delete_by_automation(self, automation_id: str, db: Optional[Session] = None) -> int: - with get_db_context(db) as db: - count = db.query(AutomationRun).filter_by(automation_id=automation_id).delete() - db.commit() - return count + async def delete_by_automation(self, automation_id: str, db: Optional[AsyncSession] = None) -> int: + async with get_async_db_context(db) as db: + result = await db.execute(delete(AutomationRun).filter_by(automation_id=automation_id)) + await db.commit() + return result.rowcount Automations = AutomationTable() diff --git a/backend/open_webui/models/channels.py b/backend/open_webui/models/channels.py index 10f2e5cad4..9b5403e6e1 100644 --- a/backend/open_webui/models/channels.py +++ b/backend/open_webui/models/channels.py @@ -4,8 +4,9 @@ import time import uuid from typing import Optional -from sqlalchemy.orm import Session -from open_webui.internal.db import Base, JSONField, get_db, get_db_context +from sqlalchemy import select, delete, update, func, case, or_, and_ +from sqlalchemy.ext.asyncio import AsyncSession +from open_webui.internal.db import Base, JSONField, get_async_db_context from open_webui.models.groups import Groups from open_webui.models.access_grants import ( AccessGrantModel, @@ -25,11 +26,7 @@ from sqlalchemy import ( Text, JSON, UniqueConstraint, - case, - cast, ) -from sqlalchemy import or_, func, select, and_, text -from sqlalchemy.sql import exists #################### # Channel DB Schema @@ -249,22 +246,22 @@ class ChannelWebhookForm(BaseModel): class ChannelTable: - def _get_access_grants(self, channel_id: str, db: Optional[Session] = None) -> list[AccessGrantModel]: - return AccessGrants.get_grants_by_resource('channel', channel_id, db=db) + async def _get_access_grants(self, channel_id: str, db: Optional[AsyncSession] = None) -> list[AccessGrantModel]: + return await AccessGrants.get_grants_by_resource('channel', channel_id, db=db) - def _to_channel_model( + async def _to_channel_model( self, channel: Channel, access_grants: Optional[list[AccessGrantModel]] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> ChannelModel: channel_data = ChannelModel.model_validate(channel).model_dump(exclude={'access_grants'}) channel_data['access_grants'] = ( - access_grants if access_grants is not None else self._get_access_grants(channel_data['id'], db=db) + access_grants if access_grants is not None else await self._get_access_grants(channel_data['id'], db=db) ) return ChannelModel.model_validate(channel_data) - def _collect_unique_user_ids( + async def _collect_unique_user_ids( self, invited_by: str, user_ids: Optional[list[str]] = None, @@ -281,7 +278,8 @@ class ChannelTable: users.add(invited_by) for group_id in group_ids or []: - users.update(Groups.get_group_user_ids_by_id(group_id)) + group_user_ids = await Groups.get_group_user_ids_by_id(group_id) + users.update(group_user_ids) return users @@ -321,10 +319,20 @@ class ChannelTable: return memberships - def insert_new_channel( - self, form_data: CreateChannelForm, user_id: str, db: Optional[Session] = None + def _has_permission(self, db, query, filter: dict, permission: str = 'read'): + return AccessGrants.has_permission_filter( + db=db, + query=query, + DocumentModel=Channel, + filter=filter, + resource_type='channel', + permission=permission, + ) + + async def insert_new_channel( + self, form_data: CreateChannelForm, user_id: str, db: Optional[AsyncSession] = None ) -> Optional[ChannelModel]: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: channel = ChannelModel( **{ **form_data.model_dump(exclude={'access_grants'}), @@ -340,7 +348,7 @@ class ChannelTable: new_channel = Channel(**channel.model_dump(exclude={'access_grants'})) if form_data.type in ['group', 'dm']: - users = self._collect_unique_user_ids( + users = await self._collect_unique_user_ids( invited_by=user_id, user_ids=form_data.user_ids, group_ids=form_data.group_ids, @@ -353,17 +361,18 @@ class ChannelTable: db.add_all(memberships) db.add(new_channel) - db.commit() - AccessGrants.set_access_grants('channel', new_channel.id, form_data.access_grants, db=db) - return self._to_channel_model(new_channel, db=db) + await db.commit() + await AccessGrants.set_access_grants('channel', new_channel.id, form_data.access_grants, db=db) + return await self._to_channel_model(new_channel, db=db) - def get_channels(self, db: Optional[Session] = None) -> list[ChannelModel]: - with get_db_context(db) as db: - channels = db.query(Channel).all() + async def get_channels(self, db: Optional[AsyncSession] = None) -> list[ChannelModel]: + async with get_async_db_context(db) as db: + result = await db.execute(select(Channel)) + channels = result.scalars().all() channel_ids = [channel.id for channel in channels] - grants_map = AccessGrants.get_grants_by_resources('channel', channel_ids, db=db) + grants_map = await AccessGrants.get_grants_by_resources('channel', channel_ids, db=db) return [ - self._to_channel_model( + await self._to_channel_model( channel, access_grants=grants_map.get(channel.id, []), db=db, @@ -371,22 +380,12 @@ class ChannelTable: for channel in channels ] - def _has_permission(self, db, query, filter: dict, permission: str = 'read'): - return AccessGrants.has_permission_filter( - db=db, - query=query, - DocumentModel=Channel, - filter=filter, - resource_type='channel', - permission=permission, - ) + async def get_channels_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[ChannelModel]: + async with get_async_db_context(db) as db: + user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id, db=db)] - def get_channels_by_user_id(self, user_id: str, db: Optional[Session] = None) -> list[ChannelModel]: - with get_db_context(db) as db: - user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id, db=db)] - - membership_channels = ( - db.query(Channel) + result = await db.execute( + select(Channel) .join(ChannelMember, Channel.id == ChannelMember.channel_id) .filter( Channel.deleted_at.is_(None), @@ -395,10 +394,10 @@ class ChannelTable: ChannelMember.user_id == user_id, ChannelMember.is_active.is_(True), ) - .all() ) + membership_channels = result.scalars().all() - query = db.query(Channel).filter( + stmt = select(Channel).filter( Channel.deleted_at.is_(None), Channel.archived_at.is_(None), or_( @@ -407,17 +406,18 @@ class ChannelTable: and_(Channel.type != 'group', Channel.type != 'dm'), ), ) - query = self._has_permission(db, query, {'user_id': user_id, 'group_ids': user_group_ids}) + stmt = self._has_permission(db, stmt, {'user_id': user_id, 'group_ids': user_group_ids}) - standard_channels = query.all() + result = await db.execute(stmt) + standard_channels = result.scalars().all() - all_channels = membership_channels + standard_channels + all_channels = list(membership_channels) + list(standard_channels) channel_ids = [c.id for c in all_channels] - grants_map = AccessGrants.get_grants_by_resources('channel', channel_ids, db=db) - return [self._to_channel_model(c, access_grants=grants_map.get(c.id, []), db=db) for c in all_channels] + grants_map = await AccessGrants.get_grants_by_resources('channel', channel_ids, db=db) + return [await self._to_channel_model(c, access_grants=grants_map.get(c.id, []), db=db) for c in all_channels] - def get_dm_channel_by_user_ids(self, user_ids: list[str], db: Optional[Session] = None) -> Optional[ChannelModel]: - with get_db_context(db) as db: + async def get_dm_channel_by_user_ids(self, user_ids: list[str], db: Optional[AsyncSession] = None) -> Optional[ChannelModel]: + async with get_async_db_context(db) as db: # Ensure uniqueness in case a list with duplicates is passed unique_user_ids = list(set(user_ids)) @@ -429,7 +429,7 @@ class ChannelTable: ) subquery = ( - db.query(ChannelMember.channel_id) + select(ChannelMember.channel_id) .group_by(ChannelMember.channel_id) # 1. Channel must have exactly len(user_ids) members .having(func.count(ChannelMember.user_id) == len(unique_user_ids)) @@ -438,33 +438,34 @@ class ChannelTable: .subquery() ) - channel = ( - db.query(Channel) + result = await db.execute( + select(Channel) .filter( - Channel.id.in_(subquery), + Channel.id.in_(select(subquery.c.channel_id)), Channel.type == 'dm', ) - .first() + .limit(1) ) + channel = result.scalars().first() - return self._to_channel_model(channel, db=db) if channel else None + return await self._to_channel_model(channel, db=db) if channel else None - def add_members_to_channel( + async def add_members_to_channel( self, channel_id: str, invited_by: str, user_ids: Optional[list[str]] = None, group_ids: Optional[list[str]] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[ChannelMemberModel]: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: # 1. Collect all user_ids including groups + inviter - requested_users = self._collect_unique_user_ids(invited_by, user_ids, group_ids) + requested_users = await self._collect_unique_user_ids(invited_by, user_ids, group_ids) - existing_users = { - row.user_id - for row in db.query(ChannelMember.user_id).filter(ChannelMember.channel_id == channel_id).all() - } + result = await db.execute( + select(ChannelMember.user_id).filter(ChannelMember.channel_id == channel_id) + ) + existing_users = {row[0] for row in result.all()} new_user_ids = requested_users - existing_users if not new_user_ids: @@ -473,59 +474,54 @@ class ChannelTable: new_memberships = self._create_membership_models(channel_id, invited_by, new_user_ids) db.add_all(new_memberships) - db.commit() + await db.commit() return [ChannelMemberModel.model_validate(membership) for membership in new_memberships] - def remove_members_from_channel( + async def remove_members_from_channel( self, channel_id: str, user_ids: list[str], - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> int: - with get_db_context(db) as db: - result = ( - db.query(ChannelMember) - .filter( + async with get_async_db_context(db) as db: + result = await db.execute( + delete(ChannelMember).filter( ChannelMember.channel_id == channel_id, ChannelMember.user_id.in_(user_ids), ) - .delete(synchronize_session=False) ) - db.commit() - return result # number of rows deleted + await db.commit() + return result.rowcount # number of rows deleted - def is_user_channel_manager(self, channel_id: str, user_id: str, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: - # Check if the user is the creator of the channel - # or has a 'manager' role in ChannelMember - channel = db.query(Channel).filter(Channel.id == channel_id).first() + async def is_user_channel_manager(self, channel_id: str, user_id: str, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: + result = await db.execute(select(Channel).filter(Channel.id == channel_id)) + channel = result.scalars().first() if channel and channel.user_id == user_id: return True - membership = ( - db.query(ChannelMember) - .filter( + result = await db.execute( + select(ChannelMember).filter( ChannelMember.channel_id == channel_id, ChannelMember.user_id == user_id, ChannelMember.is_active.is_(True), ChannelMember.role == 'manager', ) - .first() ) + membership = result.scalars().first() return membership is not None - def join_channel(self, channel_id: str, user_id: str, db: Optional[Session] = None) -> Optional[ChannelMemberModel]: - with get_db_context(db) as db: + async def join_channel(self, channel_id: str, user_id: str, db: Optional[AsyncSession] = None) -> Optional[ChannelMemberModel]: + async with get_async_db_context(db) as db: # Check if the membership already exists - existing_membership = ( - db.query(ChannelMember) - .filter( + result = await db.execute( + select(ChannelMember).filter( ChannelMember.channel_id == channel_id, ChannelMember.user_id == user_id, ) - .first() ) + existing_membership = result.scalars().first() if existing_membership: return ChannelMemberModel.model_validate(existing_membership) @@ -549,19 +545,18 @@ class ChannelTable: new_membership = ChannelMember(**channel_member.model_dump()) db.add(new_membership) - db.commit() + await db.commit() return channel_member - def leave_channel(self, channel_id: str, user_id: str, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: - membership = ( - db.query(ChannelMember) - .filter( + async def leave_channel(self, channel_id: str, user_id: str, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: + result = await db.execute( + select(ChannelMember).filter( ChannelMember.channel_id == channel_id, ChannelMember.user_id == user_id, ) - .first() ) + membership = result.scalars().first() if not membership: return False @@ -570,126 +565,127 @@ class ChannelTable: membership.left_at = int(time.time_ns()) membership.updated_at = int(time.time_ns()) - db.commit() + await db.commit() return True - def get_member_by_channel_and_user_id( - self, channel_id: str, user_id: str, db: Optional[Session] = None + async def get_member_by_channel_and_user_id( + self, channel_id: str, user_id: str, db: Optional[AsyncSession] = None ) -> Optional[ChannelMemberModel]: - with get_db_context(db) as db: - membership = ( - db.query(ChannelMember) - .filter( + async with get_async_db_context(db) as db: + result = await db.execute( + select(ChannelMember).filter( ChannelMember.channel_id == channel_id, ChannelMember.user_id == user_id, ) - .first() ) + membership = result.scalars().first() return ChannelMemberModel.model_validate(membership) if membership else None - def get_members_by_channel_id(self, channel_id: str, db: Optional[Session] = None) -> list[ChannelMemberModel]: - with get_db_context(db) as db: - memberships = db.query(ChannelMember).filter(ChannelMember.channel_id == channel_id).all() + async def get_members_by_channel_id(self, channel_id: str, db: Optional[AsyncSession] = None) -> list[ChannelMemberModel]: + async with get_async_db_context(db) as db: + result = await db.execute( + select(ChannelMember).filter(ChannelMember.channel_id == channel_id) + ) + memberships = result.scalars().all() return [ChannelMemberModel.model_validate(membership) for membership in memberships] - def pin_channel( + async def pin_channel( self, channel_id: str, user_id: str, is_pinned: bool, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> bool: - with get_db_context(db) as db: - membership = ( - db.query(ChannelMember) - .filter( + async with get_async_db_context(db) as db: + result = await db.execute( + select(ChannelMember).filter( ChannelMember.channel_id == channel_id, ChannelMember.user_id == user_id, ) - .first() ) + membership = result.scalars().first() if not membership: return False membership.is_channel_pinned = is_pinned membership.updated_at = int(time.time_ns()) - db.commit() + await db.commit() return True - def update_member_last_read_at(self, channel_id: str, user_id: str, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: - membership = ( - db.query(ChannelMember) - .filter( + async def update_member_last_read_at(self, channel_id: str, user_id: str, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: + result = await db.execute( + select(ChannelMember).filter( ChannelMember.channel_id == channel_id, ChannelMember.user_id == user_id, ) - .first() ) + membership = result.scalars().first() if not membership: return False membership.last_read_at = int(time.time_ns()) membership.updated_at = int(time.time_ns()) - db.commit() + await db.commit() return True - def update_member_active_status( + async def update_member_active_status( self, channel_id: str, user_id: str, is_active: bool, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> bool: - with get_db_context(db) as db: - membership = ( - db.query(ChannelMember) - .filter( + async with get_async_db_context(db) as db: + result = await db.execute( + select(ChannelMember).filter( ChannelMember.channel_id == channel_id, ChannelMember.user_id == user_id, ) - .first() ) + membership = result.scalars().first() if not membership: return False membership.is_active = is_active membership.updated_at = int(time.time_ns()) - db.commit() + await db.commit() return True - def is_user_channel_member(self, channel_id: str, user_id: str, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: - membership = ( - db.query(ChannelMember) - .filter( + async def is_user_channel_member(self, channel_id: str, user_id: str, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: + result = await db.execute( + select(ChannelMember).filter( ChannelMember.channel_id == channel_id, ChannelMember.user_id == user_id, ChannelMember.is_active.is_(True), - ) - .first() + ).limit(1) ) + membership = result.scalars().first() return membership is not None - def get_channel_by_id(self, id: str, db: Optional[Session] = None) -> Optional[ChannelModel]: + async def get_channel_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[ChannelModel]: try: - with get_db_context(db) as db: - channel = db.query(Channel).filter(Channel.id == id).first() - return self._to_channel_model(channel, db=db) if channel else None + async with get_async_db_context(db) as db: + result = await db.execute(select(Channel).filter(Channel.id == id)) + channel = result.scalars().first() + return await self._to_channel_model(channel, db=db) if channel else None except Exception: return None - def get_channels_by_file_id(self, file_id: str, db: Optional[Session] = None) -> list[ChannelModel]: - with get_db_context(db) as db: - channel_files = db.query(ChannelFile).filter(ChannelFile.file_id == file_id).all() + async def get_channels_by_file_id(self, file_id: str, db: Optional[AsyncSession] = None) -> list[ChannelModel]: + async with get_async_db_context(db) as db: + result = await db.execute(select(ChannelFile).filter(ChannelFile.file_id == file_id)) + channel_files = result.scalars().all() channel_ids = [cf.channel_id for cf in channel_files] - channels = db.query(Channel).filter(Channel.id.in_(channel_ids)).all() - grants_map = AccessGrants.get_grants_by_resources('channel', channel_ids, db=db) + result = await db.execute(select(Channel).filter(Channel.id.in_(channel_ids))) + channels = result.scalars().all() + grants_map = await AccessGrants.get_grants_by_resources('channel', channel_ids, db=db) return [ - self._to_channel_model( + await self._to_channel_model( channel, access_grants=grants_map.get(channel.id, []), db=db, @@ -697,123 +693,123 @@ class ChannelTable: for channel in channels ] - def get_channels_by_file_id_and_user_id( - self, file_id: str, user_id: str, db: Optional[Session] = None + async def get_channels_by_file_id_and_user_id( + self, file_id: str, user_id: str, db: Optional[AsyncSession] = None ) -> list[ChannelModel]: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: # 1. Determine which channels have this file - channel_file_rows = db.query(ChannelFile).filter(ChannelFile.file_id == file_id).all() + result = await db.execute(select(ChannelFile).filter(ChannelFile.file_id == file_id)) + channel_file_rows = result.scalars().all() channel_ids = [row.channel_id for row in channel_file_rows] if not channel_ids: return [] # 2. Load all channel rows that still exist - channels = ( - db.query(Channel) - .filter( + result = await db.execute( + select(Channel).filter( Channel.id.in_(channel_ids), Channel.deleted_at.is_(None), Channel.archived_at.is_(None), ) - .all() ) + channels = result.scalars().all() if not channels: return [] # Preload user's group membership - user_group_ids = [g.id for g in Groups.get_groups_by_member_id(user_id, db=db)] + user_group_ids = [g.id for g in await Groups.get_groups_by_member_id(user_id, db=db)] allowed_channels = [] for channel in channels: # --- Case A: group or dm => user must be an active member --- if channel.type in ['group', 'dm']: - membership = ( - db.query(ChannelMember) - .filter( + result = await db.execute( + select(ChannelMember).filter( ChannelMember.channel_id == channel.id, ChannelMember.user_id == user_id, ChannelMember.is_active.is_(True), - ) - .first() + ).limit(1) ) + membership = result.scalars().first() if membership: - allowed_channels.append(self._to_channel_model(channel, db=db)) + allowed_channels.append(await self._to_channel_model(channel, db=db)) continue # --- Case B: standard channel => rely on ACL permissions --- - query = db.query(Channel).filter(Channel.id == channel.id) + stmt = select(Channel).filter(Channel.id == channel.id) - query = self._has_permission( + stmt = self._has_permission( db, - query, + stmt, {'user_id': user_id, 'group_ids': user_group_ids}, permission='read', ) - allowed = query.first() + result = await db.execute(stmt) + allowed = result.scalars().first() if allowed: - allowed_channels.append(self._to_channel_model(allowed, db=db)) + allowed_channels.append(await self._to_channel_model(allowed, db=db)) return allowed_channels - def get_channel_by_id_and_user_id( - self, id: str, user_id: str, db: Optional[Session] = None + async def get_channel_by_id_and_user_id( + self, id: str, user_id: str, db: Optional[AsyncSession] = None ) -> Optional[ChannelModel]: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: # Fetch the channel - channel: Channel = ( - db.query(Channel) - .filter( + result = await db.execute( + select(Channel).filter( Channel.id == id, Channel.deleted_at.is_(None), Channel.archived_at.is_(None), ) - .first() ) + channel = result.scalars().first() if not channel: return None # If the channel is a group or dm, read access requires membership (active) if channel.type in ['group', 'dm']: - membership = ( - db.query(ChannelMember) - .filter( + result = await db.execute( + select(ChannelMember).filter( ChannelMember.channel_id == id, ChannelMember.user_id == user_id, ChannelMember.is_active.is_(True), - ) - .first() + ).limit(1) ) + membership = result.scalars().first() if membership: - return self._to_channel_model(channel, db=db) + return await self._to_channel_model(channel, db=db) else: return None # For channels that are NOT group/dm, fall back to ACL-based read access - query = db.query(Channel).filter(Channel.id == id) + stmt = select(Channel).filter(Channel.id == id) # Determine user groups - user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id, db=db)] + user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id, db=db)] # Apply ACL rules - query = self._has_permission( + stmt = self._has_permission( db, - query, + stmt, {'user_id': user_id, 'group_ids': user_group_ids}, permission='read', ) - channel_allowed = query.first() - return self._to_channel_model(channel_allowed, db=db) if channel_allowed else None + result = await db.execute(stmt) + channel_allowed = result.scalars().first() + return await self._to_channel_model(channel_allowed, db=db) if channel_allowed else None - def update_channel_by_id( - self, id: str, form_data: ChannelForm, db: Optional[Session] = None + async def update_channel_by_id( + self, id: str, form_data: ChannelForm, db: Optional[AsyncSession] = None ) -> Optional[ChannelModel]: - with get_db_context(db) as db: - channel = db.query(Channel).filter(Channel.id == id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(Channel).filter(Channel.id == id)) + channel = result.scalars().first() if not channel: return None @@ -825,16 +821,16 @@ class ChannelTable: channel.meta = form_data.meta if form_data.access_grants is not None: - AccessGrants.set_access_grants('channel', id, form_data.access_grants, db=db) + await AccessGrants.set_access_grants('channel', id, form_data.access_grants, db=db) channel.updated_at = int(time.time_ns()) - db.commit() - return self._to_channel_model(channel, db=db) if channel else None + await db.commit() + return await self._to_channel_model(channel, db=db) if channel else None - def add_file_to_channel_by_id( - self, channel_id: str, file_id: str, user_id: str, db: Optional[Session] = None + async def add_file_to_channel_by_id( + self, channel_id: str, file_id: str, user_id: str, db: Optional[AsyncSession] = None ) -> Optional[ChannelFileModel]: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: channel_file = ChannelFileModel( **{ 'id': str(uuid.uuid4()), @@ -849,8 +845,8 @@ class ChannelTable: try: result = ChannelFile(**channel_file.model_dump()) db.add(result) - db.commit() - db.refresh(result) + await db.commit() + await db.refresh(result) if result: return ChannelFileModel.model_validate(result) else: @@ -858,55 +854,58 @@ class ChannelTable: except Exception: return None - def set_file_message_id_in_channel_by_id( + async def set_file_message_id_in_channel_by_id( self, channel_id: str, file_id: str, message_id: str, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> bool: try: - with get_db_context(db) as db: - channel_file = db.query(ChannelFile).filter_by(channel_id=channel_id, file_id=file_id).first() + async with get_async_db_context(db) as db: + result = await db.execute( + select(ChannelFile).filter_by(channel_id=channel_id, file_id=file_id) + ) + channel_file = result.scalars().first() if not channel_file: return False channel_file.message_id = message_id channel_file.updated_at = int(time.time()) - db.commit() + await db.commit() return True except Exception: return False - def remove_file_from_channel_by_id(self, channel_id: str, file_id: str, db: Optional[Session] = None) -> bool: + async def remove_file_from_channel_by_id(self, channel_id: str, file_id: str, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: - db.query(ChannelFile).filter_by(channel_id=channel_id, file_id=file_id).delete() - db.commit() + async with get_async_db_context(db) as db: + await db.execute(delete(ChannelFile).filter_by(channel_id=channel_id, file_id=file_id)) + await db.commit() return True except Exception: return False - def delete_channel_by_id(self, id: str, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: - AccessGrants.revoke_all_access('channel', id, db=db) - db.query(Channel).filter(Channel.id == id).delete() - db.commit() + async def delete_channel_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: + await AccessGrants.revoke_all_access('channel', id, db=db) + await db.execute(delete(Channel).filter(Channel.id == id)) + await db.commit() return True #################### # Webhook Methods #################### - def insert_webhook( + async def insert_webhook( self, channel_id: str, user_id: str, form_data: ChannelWebhookForm, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[ChannelWebhookModel]: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: webhook = ChannelWebhookModel( id=str(uuid.uuid4()), channel_id=channel_id, @@ -919,63 +918,66 @@ class ChannelTable: updated_at=int(time.time_ns()), ) db.add(ChannelWebhook(**webhook.model_dump())) - db.commit() + await db.commit() return webhook - def get_webhooks_by_channel_id(self, channel_id: str, db: Optional[Session] = None) -> list[ChannelWebhookModel]: - with get_db_context(db) as db: - webhooks = db.query(ChannelWebhook).filter(ChannelWebhook.channel_id == channel_id).all() + async def get_webhooks_by_channel_id(self, channel_id: str, db: Optional[AsyncSession] = None) -> list[ChannelWebhookModel]: + async with get_async_db_context(db) as db: + result = await db.execute(select(ChannelWebhook).filter(ChannelWebhook.channel_id == channel_id)) + webhooks = result.scalars().all() return [ChannelWebhookModel.model_validate(w) for w in webhooks] - def get_webhook_by_id(self, webhook_id: str, db: Optional[Session] = None) -> Optional[ChannelWebhookModel]: - with get_db_context(db) as db: - webhook = db.query(ChannelWebhook).filter(ChannelWebhook.id == webhook_id).first() + async def get_webhook_by_id(self, webhook_id: str, db: Optional[AsyncSession] = None) -> Optional[ChannelWebhookModel]: + async with get_async_db_context(db) as db: + result = await db.execute(select(ChannelWebhook).filter(ChannelWebhook.id == webhook_id)) + webhook = result.scalars().first() return ChannelWebhookModel.model_validate(webhook) if webhook else None - def get_webhook_by_id_and_token( - self, webhook_id: str, token: str, db: Optional[Session] = None + async def get_webhook_by_id_and_token( + self, webhook_id: str, token: str, db: Optional[AsyncSession] = None ) -> Optional[ChannelWebhookModel]: - with get_db_context(db) as db: - webhook = ( - db.query(ChannelWebhook) - .filter( + async with get_async_db_context(db) as db: + result = await db.execute( + select(ChannelWebhook).filter( ChannelWebhook.id == webhook_id, ChannelWebhook.token == token, ) - .first() ) + webhook = result.scalars().first() return ChannelWebhookModel.model_validate(webhook) if webhook else None - def update_webhook_by_id( + async def update_webhook_by_id( self, webhook_id: str, form_data: ChannelWebhookForm, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[ChannelWebhookModel]: - with get_db_context(db) as db: - webhook = db.query(ChannelWebhook).filter(ChannelWebhook.id == webhook_id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(ChannelWebhook).filter(ChannelWebhook.id == webhook_id)) + webhook = result.scalars().first() if not webhook: return None webhook.name = form_data.name webhook.profile_image_url = form_data.profile_image_url webhook.updated_at = int(time.time_ns()) - db.commit() + await db.commit() return ChannelWebhookModel.model_validate(webhook) - def update_webhook_last_used_at(self, webhook_id: str, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: - webhook = db.query(ChannelWebhook).filter(ChannelWebhook.id == webhook_id).first() + async def update_webhook_last_used_at(self, webhook_id: str, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: + result = await db.execute(select(ChannelWebhook).filter(ChannelWebhook.id == webhook_id)) + webhook = result.scalars().first() if not webhook: return False webhook.last_used_at = int(time.time_ns()) - db.commit() + await db.commit() return True - def delete_webhook_by_id(self, webhook_id: str, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: - result = db.query(ChannelWebhook).filter(ChannelWebhook.id == webhook_id).delete() - db.commit() - return result > 0 + async def delete_webhook_by_id(self, webhook_id: str, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: + result = await db.execute(delete(ChannelWebhook).filter(ChannelWebhook.id == webhook_id)) + await db.commit() + return result.rowcount > 0 Channels = ChannelTable() diff --git a/backend/open_webui/models/chat_messages.py b/backend/open_webui/models/chat_messages.py index b37c04037e..087662ff7c 100644 --- a/backend/open_webui/models/chat_messages.py +++ b/backend/open_webui/models/chat_messages.py @@ -3,8 +3,9 @@ import time import uuid from typing import Any, Optional -from sqlalchemy.orm import Session -from open_webui.internal.db import Base, get_db_context +from sqlalchemy import select, delete, func, cast, Integer +from sqlalchemy.ext.asyncio import AsyncSession +from open_webui.internal.db import Base, get_async_db_context from open_webui.utils.response import normalize_usage from pydantic import BaseModel, ConfigDict @@ -16,7 +17,6 @@ from sqlalchemy import ( Text, JSON, Index, - func, ) #################### @@ -129,23 +129,23 @@ class ChatMessageModel(BaseModel): class ChatMessageTable: - def upsert_message( + async def upsert_message( self, message_id: str, chat_id: str, user_id: str, data: dict, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[ChatMessageModel]: """Insert or update a chat message.""" - with get_db_context(db) as db: + async with get_async_db_context(db) as db: now = int(time.time()) timestamp = data.get('timestamp', now) # Use composite ID: {chat_id}-{message_id} composite_id = f'{chat_id}-{message_id}' - existing = db.get(ChatMessage, composite_id) + existing = await db.get(ChatMessage, composite_id) if existing: # Update existing if 'role' in data: @@ -178,8 +178,8 @@ class ChatMessageTable: # from accidentally clearing the primary response's token counts existing.usage = {**(existing.usage or {}), **usage} existing.updated_at = now - db.commit() - db.refresh(existing) + await db.commit() + await db.refresh(existing) return ChatMessageModel.model_validate(existing) else: # Insert new @@ -205,143 +205,155 @@ class ChatMessageTable: updated_at=now, ) db.add(message) - db.commit() - db.refresh(message) + await db.commit() + await db.refresh(message) return ChatMessageModel.model_validate(message) - def get_message_by_id(self, id: str, db: Optional[Session] = None) -> Optional[ChatMessageModel]: - with get_db_context(db) as db: - message = db.get(ChatMessage, id) + async def get_message_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[ChatMessageModel]: + async with get_async_db_context(db) as db: + message = await db.get(ChatMessage, id) return ChatMessageModel.model_validate(message) if message else None - def get_messages_by_chat_id(self, chat_id: str, db: Optional[Session] = None) -> list[ChatMessageModel]: - with get_db_context(db) as db: - messages = db.query(ChatMessage).filter_by(chat_id=chat_id).order_by(ChatMessage.created_at.asc()).all() + async def get_messages_by_chat_id(self, chat_id: str, db: Optional[AsyncSession] = None) -> list[ChatMessageModel]: + async with get_async_db_context(db) as db: + result = await db.execute( + select(ChatMessage).filter_by(chat_id=chat_id).order_by(ChatMessage.created_at.asc()) + ) + messages = result.scalars().all() return [ChatMessageModel.model_validate(message) for message in messages] - def get_messages_by_user_id( + async def get_messages_by_user_id( self, user_id: str, skip: int = 0, limit: int = 50, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[ChatMessageModel]: - with get_db_context(db) as db: - messages = ( - db.query(ChatMessage) + async with get_async_db_context(db) as db: + result = await db.execute( + select(ChatMessage) .filter_by(user_id=user_id) .order_by(ChatMessage.created_at.desc()) .offset(skip) .limit(limit) - .all() ) + messages = result.scalars().all() return [ChatMessageModel.model_validate(message) for message in messages] - def get_messages_by_model_id( + async def get_messages_by_model_id( self, model_id: str, start_date: Optional[int] = None, end_date: Optional[int] = None, skip: int = 0, limit: int = 100, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[ChatMessageModel]: - with get_db_context(db) as db: - query = db.query(ChatMessage).filter_by(model_id=model_id) + async with get_async_db_context(db) as db: + stmt = select(ChatMessage).filter_by(model_id=model_id) if start_date: - query = query.filter(ChatMessage.created_at >= start_date) + stmt = stmt.filter(ChatMessage.created_at >= start_date) if end_date: - query = query.filter(ChatMessage.created_at <= end_date) - messages = query.order_by(ChatMessage.created_at.desc()).offset(skip).limit(limit).all() + stmt = stmt.filter(ChatMessage.created_at <= end_date) + stmt = stmt.order_by(ChatMessage.created_at.desc()).offset(skip).limit(limit) + result = await db.execute(stmt) + messages = result.scalars().all() return [ChatMessageModel.model_validate(message) for message in messages] - def get_chat_ids_by_model_id( + async def get_chat_ids_by_model_id( self, model_id: str, start_date: Optional[int] = None, end_date: Optional[int] = None, skip: int = 0, limit: int = 50, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[str]: """Get distinct chat_ids that used a specific model.""" - with get_db_context(db) as db: - query = db.query( - ChatMessage.chat_id, - func.max(ChatMessage.created_at).label('last_message_at'), - ).filter(ChatMessage.model_id == model_id) + async with get_async_db_context(db) as db: + stmt = ( + select( + ChatMessage.chat_id, + func.max(ChatMessage.created_at).label('last_message_at'), + ) + .filter(ChatMessage.model_id == model_id) + ) if start_date: - query = query.filter(ChatMessage.created_at >= start_date) + stmt = stmt.filter(ChatMessage.created_at >= start_date) if end_date: - query = query.filter(ChatMessage.created_at <= end_date) + stmt = stmt.filter(ChatMessage.created_at <= end_date) # Group by chat_id and order by most recent message in each chat # Secondary sort on chat_id ensures deterministic pagination - # (prevents duplicates across pages when timestamps tie) - chat_ids = ( - query.group_by(ChatMessage.chat_id) + stmt = ( + stmt.group_by(ChatMessage.chat_id) .order_by(func.max(ChatMessage.created_at).desc(), ChatMessage.chat_id) .offset(skip) .limit(limit) - .all() ) + result = await db.execute(stmt) + chat_ids = result.all() return [chat_id for chat_id, _ in chat_ids] - def delete_messages_by_chat_id(self, chat_id: str, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: - db.query(ChatMessage).filter_by(chat_id=chat_id).delete() - db.commit() + async def delete_messages_by_chat_id(self, chat_id: str, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: + await db.execute(delete(ChatMessage).filter_by(chat_id=chat_id)) + await db.commit() return True # Analytics methods - def get_message_count_by_model( + async def get_message_count_by_model( self, start_date: Optional[int] = None, end_date: Optional[int] = None, group_id: Optional[str] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> dict[str, int]: - with get_db_context(db) as db: - from sqlalchemy import func + async with get_async_db_context(db) as db: from open_webui.models.groups import GroupMember - query = db.query(ChatMessage.model_id, func.count(ChatMessage.id).label('count')).filter( - ChatMessage.role == 'assistant', - ChatMessage.model_id.isnot(None), - ~ChatMessage.user_id.like('shared-%'), + stmt = ( + select(ChatMessage.model_id, func.count(ChatMessage.id).label('count')) + .filter( + ChatMessage.role == 'assistant', + ChatMessage.model_id.isnot(None), + ~ChatMessage.user_id.like('shared-%'), + ) ) if start_date: - query = query.filter(ChatMessage.created_at >= start_date) + stmt = stmt.filter(ChatMessage.created_at >= start_date) if end_date: - query = query.filter(ChatMessage.created_at <= end_date) + stmt = stmt.filter(ChatMessage.created_at <= end_date) if group_id: - group_users = db.query(GroupMember.user_id).filter(GroupMember.group_id == group_id).subquery() - query = query.filter(ChatMessage.user_id.in_(group_users)) + group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery() + stmt = stmt.filter(ChatMessage.user_id.in_(group_users)) - results = query.group_by(ChatMessage.model_id).all() - return {row.model_id: row.count for row in results} + stmt = stmt.group_by(ChatMessage.model_id) + result = await db.execute(stmt) + return {row.model_id: row.count for row in result.all()} - def get_token_usage_by_model( + async def get_token_usage_by_model( self, start_date: Optional[int] = None, end_date: Optional[int] = None, group_id: Optional[str] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> dict[str, dict]: """Aggregate token usage by model using database-level aggregation.""" - with get_db_context(db) as db: - from sqlalchemy import func, cast, Integer + async with get_async_db_context(db) as db: from open_webui.models.groups import GroupMember - dialect = db.bind.dialect.name + # We need the dialect to determine JSON extraction syntax + # For async sessions, access via get_bind() + bind = await db.connection() + dialect = bind.dialect.name if dialect == 'sqlite': input_tokens = cast(func.json_extract(ChatMessage.usage, '$.input_tokens'), Integer) output_tokens = cast(func.json_extract(ChatMessage.usage, '$.output_tokens'), Integer) elif dialect == 'postgresql': - # Use json_extract_path_text for PostgreSQL JSON columns input_tokens = cast( func.json_extract_path_text(ChatMessage.usage, 'input_tokens'), Integer, @@ -353,27 +365,31 @@ class ChatMessageTable: else: raise NotImplementedError(f'Unsupported dialect: {dialect}') - query = db.query( - ChatMessage.model_id, - func.coalesce(func.sum(input_tokens), 0).label('input_tokens'), - func.coalesce(func.sum(output_tokens), 0).label('output_tokens'), - func.count(ChatMessage.id).label('message_count'), - ).filter( - ChatMessage.role == 'assistant', - ChatMessage.model_id.isnot(None), - ChatMessage.usage.isnot(None), - ~ChatMessage.user_id.like('shared-%'), + stmt = ( + select( + ChatMessage.model_id, + func.coalesce(func.sum(input_tokens), 0).label('input_tokens'), + func.coalesce(func.sum(output_tokens), 0).label('output_tokens'), + func.count(ChatMessage.id).label('message_count'), + ) + .filter( + ChatMessage.role == 'assistant', + ChatMessage.model_id.isnot(None), + ChatMessage.usage.isnot(None), + ~ChatMessage.user_id.like('shared-%'), + ) ) if start_date: - query = query.filter(ChatMessage.created_at >= start_date) + stmt = stmt.filter(ChatMessage.created_at >= start_date) if end_date: - query = query.filter(ChatMessage.created_at <= end_date) + stmt = stmt.filter(ChatMessage.created_at <= end_date) if group_id: - group_users = db.query(GroupMember.user_id).filter(GroupMember.group_id == group_id).subquery() - query = query.filter(ChatMessage.user_id.in_(group_users)) + group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery() + stmt = stmt.filter(ChatMessage.user_id.in_(group_users)) - results = query.group_by(ChatMessage.model_id).all() + stmt = stmt.group_by(ChatMessage.model_id) + result = await db.execute(stmt) return { row.model_id: { @@ -382,28 +398,27 @@ class ChatMessageTable: 'total_tokens': row.input_tokens + row.output_tokens, 'message_count': row.message_count, } - for row in results + for row in result.all() } - def get_token_usage_by_user( + async def get_token_usage_by_user( self, start_date: Optional[int] = None, end_date: Optional[int] = None, group_id: Optional[str] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> dict[str, dict]: """Aggregate token usage by user using database-level aggregation.""" - with get_db_context(db) as db: - from sqlalchemy import func, cast, Integer + async with get_async_db_context(db) as db: from open_webui.models.groups import GroupMember - dialect = db.bind.dialect.name + bind = await db.connection() + dialect = bind.dialect.name if dialect == 'sqlite': input_tokens = cast(func.json_extract(ChatMessage.usage, '$.input_tokens'), Integer) output_tokens = cast(func.json_extract(ChatMessage.usage, '$.output_tokens'), Integer) elif dialect == 'postgresql': - # Use json_extract_path_text for PostgreSQL JSON columns input_tokens = cast( func.json_extract_path_text(ChatMessage.usage, 'input_tokens'), Integer, @@ -415,27 +430,31 @@ class ChatMessageTable: else: raise NotImplementedError(f'Unsupported dialect: {dialect}') - query = db.query( - ChatMessage.user_id, - func.coalesce(func.sum(input_tokens), 0).label('input_tokens'), - func.coalesce(func.sum(output_tokens), 0).label('output_tokens'), - func.count(ChatMessage.id).label('message_count'), - ).filter( - ChatMessage.role == 'assistant', - ChatMessage.user_id.isnot(None), - ChatMessage.usage.isnot(None), - ~ChatMessage.user_id.like('shared-%'), + stmt = ( + select( + ChatMessage.user_id, + func.coalesce(func.sum(input_tokens), 0).label('input_tokens'), + func.coalesce(func.sum(output_tokens), 0).label('output_tokens'), + func.count(ChatMessage.id).label('message_count'), + ) + .filter( + ChatMessage.role == 'assistant', + ChatMessage.user_id.isnot(None), + ChatMessage.usage.isnot(None), + ~ChatMessage.user_id.like('shared-%'), + ) ) if start_date: - query = query.filter(ChatMessage.created_at >= start_date) + stmt = stmt.filter(ChatMessage.created_at >= start_date) if end_date: - query = query.filter(ChatMessage.created_at <= end_date) + stmt = stmt.filter(ChatMessage.created_at <= end_date) if group_id: - group_users = db.query(GroupMember.user_id).filter(GroupMember.group_id == group_id).subquery() - query = query.filter(ChatMessage.user_id.in_(group_users)) + group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery() + stmt = stmt.filter(ChatMessage.user_id.in_(group_users)) - results = query.group_by(ChatMessage.user_id).all() + stmt = stmt.group_by(ChatMessage.user_id) + result = await db.execute(stmt) return { row.user_id: { @@ -444,88 +463,94 @@ class ChatMessageTable: 'total_tokens': row.input_tokens + row.output_tokens, 'message_count': row.message_count, } - for row in results + for row in result.all() } - def get_message_count_by_user( + async def get_message_count_by_user( self, start_date: Optional[int] = None, end_date: Optional[int] = None, group_id: Optional[str] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> dict[str, int]: - with get_db_context(db) as db: - from sqlalchemy import func + async with get_async_db_context(db) as db: from open_webui.models.groups import GroupMember - query = db.query(ChatMessage.user_id, func.count(ChatMessage.id).label('count')).filter( - ~ChatMessage.user_id.like('shared-%') + stmt = ( + select(ChatMessage.user_id, func.count(ChatMessage.id).label('count')) + .filter(~ChatMessage.user_id.like('shared-%')) ) if start_date: - query = query.filter(ChatMessage.created_at >= start_date) + stmt = stmt.filter(ChatMessage.created_at >= start_date) if end_date: - query = query.filter(ChatMessage.created_at <= end_date) + stmt = stmt.filter(ChatMessage.created_at <= end_date) if group_id: - group_users = db.query(GroupMember.user_id).filter(GroupMember.group_id == group_id).subquery() - query = query.filter(ChatMessage.user_id.in_(group_users)) + group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery() + stmt = stmt.filter(ChatMessage.user_id.in_(group_users)) - results = query.group_by(ChatMessage.user_id).all() - return {row.user_id: row.count for row in results} + stmt = stmt.group_by(ChatMessage.user_id) + result = await db.execute(stmt) + return {row.user_id: row.count for row in result.all()} - def get_message_count_by_chat( + async def get_message_count_by_chat( self, start_date: Optional[int] = None, end_date: Optional[int] = None, group_id: Optional[str] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> dict[str, int]: - with get_db_context(db) as db: - from sqlalchemy import func + async with get_async_db_context(db) as db: from open_webui.models.groups import GroupMember - query = db.query(ChatMessage.chat_id, func.count(ChatMessage.id).label('count')).filter( - ~ChatMessage.user_id.like('shared-%') + stmt = ( + select(ChatMessage.chat_id, func.count(ChatMessage.id).label('count')) + .filter(~ChatMessage.user_id.like('shared-%')) ) if start_date: - query = query.filter(ChatMessage.created_at >= start_date) + stmt = stmt.filter(ChatMessage.created_at >= start_date) if end_date: - query = query.filter(ChatMessage.created_at <= end_date) + stmt = stmt.filter(ChatMessage.created_at <= end_date) if group_id: - group_users = db.query(GroupMember.user_id).filter(GroupMember.group_id == group_id).subquery() - query = query.filter(ChatMessage.user_id.in_(group_users)) + group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery() + stmt = stmt.filter(ChatMessage.user_id.in_(group_users)) - results = query.group_by(ChatMessage.chat_id).all() - return {row.chat_id: row.count for row in results} + stmt = stmt.group_by(ChatMessage.chat_id) + result = await db.execute(stmt) + return {row.chat_id: row.count for row in result.all()} - def get_daily_message_counts_by_model( + async def get_daily_message_counts_by_model( self, start_date: Optional[int] = None, end_date: Optional[int] = None, group_id: Optional[str] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> dict[str, dict[str, int]]: """Get message counts grouped by day and model.""" - with get_db_context(db) as db: + async with get_async_db_context(db) as db: from datetime import datetime, timedelta from open_webui.models.groups import GroupMember - query = db.query(ChatMessage.created_at, ChatMessage.model_id).filter( - ChatMessage.role == 'assistant', - ChatMessage.model_id.isnot(None), - ~ChatMessage.user_id.like('shared-%'), + stmt = ( + select(ChatMessage.created_at, ChatMessage.model_id) + .filter( + ChatMessage.role == 'assistant', + ChatMessage.model_id.isnot(None), + ~ChatMessage.user_id.like('shared-%'), + ) ) if start_date: - query = query.filter(ChatMessage.created_at >= start_date) + stmt = stmt.filter(ChatMessage.created_at >= start_date) if end_date: - query = query.filter(ChatMessage.created_at <= end_date) + stmt = stmt.filter(ChatMessage.created_at <= end_date) if group_id: - group_users = db.query(GroupMember.user_id).filter(GroupMember.group_id == group_id).subquery() - query = query.filter(ChatMessage.user_id.in_(group_users)) + group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery() + stmt = stmt.filter(ChatMessage.user_id.in_(group_users)) - results = query.all() + result = await db.execute(stmt) + results = result.all() # Group by date -> model -> count daily_counts: dict[str, dict[str, int]] = {} @@ -547,28 +572,32 @@ class ChatMessageTable: return daily_counts - def get_hourly_message_counts_by_model( + async def get_hourly_message_counts_by_model( self, start_date: Optional[int] = None, end_date: Optional[int] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> dict[str, dict[str, int]]: """Get message counts grouped by hour and model.""" - with get_db_context(db) as db: + async with get_async_db_context(db) as db: from datetime import datetime, timedelta - query = db.query(ChatMessage.created_at, ChatMessage.model_id).filter( - ChatMessage.role == 'assistant', - ChatMessage.model_id.isnot(None), - ~ChatMessage.user_id.like('shared-%'), + stmt = ( + select(ChatMessage.created_at, ChatMessage.model_id) + .filter( + ChatMessage.role == 'assistant', + ChatMessage.model_id.isnot(None), + ~ChatMessage.user_id.like('shared-%'), + ) ) if start_date: - query = query.filter(ChatMessage.created_at >= start_date) + stmt = stmt.filter(ChatMessage.created_at >= start_date) if end_date: - query = query.filter(ChatMessage.created_at <= end_date) + stmt = stmt.filter(ChatMessage.created_at <= end_date) - results = query.all() + result = await db.execute(stmt) + results = result.all() # Group by hour -> model -> count hourly_counts: dict[str, dict[str, int]] = {} diff --git a/backend/open_webui/models/chats.py b/backend/open_webui/models/chats.py index 5f53e741e4..77bc0a5614 100644 --- a/backend/open_webui/models/chats.py +++ b/backend/open_webui/models/chats.py @@ -4,8 +4,11 @@ import time import uuid from typing import Optional -from sqlalchemy.orm import Session -from open_webui.internal.db import Base, JSONField, get_db, get_db_context +from sqlalchemy import select, delete, update, func, or_, and_, text +from sqlalchemy.ext.asyncio import AsyncSession +from sqlalchemy.sql import exists +from sqlalchemy.sql.expression import bindparam +from open_webui.internal.db import Base, JSONField, get_async_db_context from open_webui.models.tags import TagModel, Tag, Tags from open_webui.models.folders import Folders from open_webui.models.chat_messages import ChatMessage, ChatMessages @@ -24,9 +27,6 @@ from sqlalchemy import ( Index, UniqueConstraint, ) -from sqlalchemy import or_, func, select, and_, text -from sqlalchemy.sql import exists -from sqlalchemy.sql.expression import bindparam #################### # Chat DB Schema @@ -62,15 +62,10 @@ class Chat(Base): __table_args__ = ( # Performance indexes for common queries - # WHERE folder_id = ... Index('folder_id_idx', 'folder_id'), - # WHERE user_id = ... AND pinned = ... Index('user_id_pinned_idx', 'user_id', 'pinned'), - # WHERE user_id = ... AND archived = ... Index('user_id_archived_idx', 'user_id', 'archived'), - # WHERE user_id = ... ORDER BY updated_at DESC Index('updated_at_user_id_idx', 'updated_at', 'user_id'), - # WHERE folder_id = ... AND user_id = ... Index('folder_id_user_id_idx', 'folder_id', 'user_id'), ) @@ -297,8 +292,8 @@ class ChatTable: return changed - def insert_new_chat(self, user_id: str, form_data: ChatForm, db: Optional[Session] = None) -> Optional[ChatModel]: - with get_db_context(db) as db: + async def insert_new_chat(self, user_id: str, form_data: ChatForm, db: Optional[AsyncSession] = None) -> Optional[ChatModel]: + async with get_async_db_context(db) as db: id = str(uuid.uuid4()) chat = ChatModel( **{ @@ -316,8 +311,8 @@ class ChatTable: chat_item = Chat(**chat.model_dump()) db.add(chat_item) - db.commit() - db.refresh(chat_item) + await db.commit() + await db.refresh(chat_item) # Dual-write initial messages to chat_message table try: @@ -325,7 +320,7 @@ class ChatTable: messages = history.get('messages', {}) for message_id, message in messages.items(): if isinstance(message, dict) and message.get('role'): - ChatMessages.upsert_message( + await ChatMessages.upsert_message( message_id=message_id, chat_id=id, user_id=user_id, @@ -353,13 +348,13 @@ class ChatTable: ) return chat - def import_chats( + async def import_chats( self, user_id: str, chat_import_forms: list[ChatImportForm], - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[ChatModel]: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: chats = [] for form_data in chat_import_forms: @@ -367,7 +362,7 @@ class ChatTable: chats.append(Chat(**chat.model_dump())) db.add_all(chats) - db.commit() + await db.commit() # Dual-write messages to chat_message table try: @@ -376,7 +371,7 @@ class ChatTable: messages = history.get('messages', {}) for message_id, message in messages.items(): if isinstance(message, dict) and message.get('role'): - ChatMessages.upsert_message( + await ChatMessages.upsert_message( message_id=message_id, chat_id=chat_obj.id, user_id=user_id, @@ -387,53 +382,53 @@ class ChatTable: return [ChatModel.model_validate(chat) for chat in chats] - def update_chat_by_id(self, id: str, chat: dict, db: Optional[Session] = None) -> Optional[ChatModel]: + async def update_chat_by_id(self, id: str, chat: dict, db: Optional[AsyncSession] = None) -> Optional[ChatModel]: try: - with get_db_context(db) as db: - chat_item = db.get(Chat, id) + async with get_async_db_context(db) as db: + chat_item = await db.get(Chat, id) chat_item.chat = self._clean_null_bytes(chat) chat_item.title = self._clean_null_bytes(chat['title']) if 'title' in chat else 'New Chat' chat_item.updated_at = int(time.time()) - db.commit() - db.refresh(chat_item) + await db.commit() + await db.refresh(chat_item) return ChatModel.model_validate(chat_item) except Exception: return None - def update_chat_last_read_at_by_id(self, id: str, user_id: str, db: Optional[Session] = None) -> bool: + async def update_chat_last_read_at_by_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: - chat = db.get(Chat, id) + async with get_async_db_context(db) as db: + chat = await db.get(Chat, id) if chat and chat.user_id == user_id: chat.last_read_at = int(time.time()) - db.commit() + await db.commit() return True return False except Exception: return False - def update_chat_title_by_id(self, id: str, title: str) -> Optional[ChatModel]: + async def update_chat_title_by_id(self, id: str, title: str) -> Optional[ChatModel]: try: - with get_db_context() as db: - chat_item = db.get(Chat, id) + async with get_async_db_context() as db: + chat_item = await db.get(Chat, id) if chat_item is None: return None clean_title = self._clean_null_bytes(title) chat_item.title = clean_title chat_item.chat = {**(chat_item.chat or {}), 'title': clean_title} chat_item.updated_at = int(time.time()) - db.commit() - db.refresh(chat_item) + await db.commit() + await db.refresh(chat_item) return ChatModel.model_validate(chat_item) except Exception: return None - def update_chat_tags_by_id(self, id: str, tags: list[str], user) -> Optional[ChatModel]: - with get_db_context() as db: - chat = db.get(Chat, id) + async def update_chat_tags_by_id(self, id: str, tags: list[str], user) -> Optional[ChatModel]: + async with get_async_db_context() as db: + chat = await db.get(Chat, id) if chat is None: return None @@ -443,44 +438,45 @@ class ChatTable: # Single meta update chat.meta = {**chat.meta, 'tags': new_tag_ids} - db.commit() - db.refresh(chat) + await db.commit() + await db.refresh(chat) # Batch-create any missing tag rows - Tags.ensure_tags_exist(new_tags, user.id, db=db) + await Tags.ensure_tags_exist(new_tags, user.id, db=db) # Clean up orphaned old tags in one query removed = set(old_tags) - set(new_tag_ids) if removed: - self.delete_orphan_tags_for_user(list(removed), user.id, db=db) + await self.delete_orphan_tags_for_user(list(removed), user.id, db=db) return ChatModel.model_validate(chat) - def get_chat_title_by_id(self, id: str) -> Optional[str]: - with get_db_context() as db: - result = db.query(Chat.title).filter_by(id=id).first() - if result is None: + async def get_chat_title_by_id(self, id: str) -> Optional[str]: + async with get_async_db_context() as db: + result = await db.execute(select(Chat.title).filter_by(id=id)) + row = result.first() + if row is None: return None - return result[0] or 'New Chat' + return row[0] or 'New Chat' - def get_messages_map_by_chat_id(self, id: str) -> Optional[dict]: - chat = self.get_chat_by_id(id) + async def get_messages_map_by_chat_id(self, id: str) -> Optional[dict]: + chat = await self.get_chat_by_id(id) if chat is None: return None return chat.chat.get('history', {}).get('messages', {}) or {} - def get_message_by_id_and_message_id(self, id: str, message_id: str) -> Optional[dict]: - chat = self.get_chat_by_id(id) + async def get_message_by_id_and_message_id(self, id: str, message_id: str) -> Optional[dict]: + chat = await self.get_chat_by_id(id) if chat is None: return None return chat.chat.get('history', {}).get('messages', {}).get(message_id, {}) - def upsert_message_to_chat_by_id_and_message_id( + async def upsert_message_to_chat_by_id_and_message_id( self, id: str, message_id: str, message: dict ) -> Optional[ChatModel]: - chat = self.get_chat_by_id(id) + chat = await self.get_chat_by_id(id) if chat is None: return None @@ -506,7 +502,7 @@ class ChatTable: # Dual-write to chat_message table try: - ChatMessages.upsert_message( + await ChatMessages.upsert_message( message_id=message_id, chat_id=id, user_id=user_id, @@ -515,12 +511,12 @@ class ChatTable: except Exception as e: log.warning(f'Failed to write to chat_message table: {e}') - return self.update_chat_by_id(id, chat) + return await self.update_chat_by_id(id, chat) - def add_message_status_to_chat_by_id_and_message_id( + async def add_message_status_to_chat_by_id_and_message_id( self, id: str, message_id: str, status: dict ) -> Optional[ChatModel]: - chat = self.get_chat_by_id(id) + chat = await self.get_chat_by_id(id) if chat is None: return None @@ -533,11 +529,11 @@ class ChatTable: history['messages'][message_id]['statusHistory'] = status_history chat['history'] = history - return self.update_chat_by_id(id, chat) + return await self.update_chat_by_id(id, chat) - def add_message_files_by_id_and_message_id(self, id: str, message_id: str, files: list[dict]) -> list[dict]: - with get_db_context() as db: - chat = self.get_chat_by_id(id, db=db) + async def add_message_files_by_id_and_message_id(self, id: str, message_id: str, files: list[dict]) -> list[dict]: + async with get_async_db_context() as db: + chat = await self.get_chat_by_id(id, db=db) if chat is None: return None @@ -552,19 +548,19 @@ class ChatTable: history['messages'][message_id]['files'] = message_files chat['history'] = history - self.update_chat_by_id(id, chat, db=db) + await self.update_chat_by_id(id, chat, db=db) return message_files - def insert_shared_chat_by_chat_id(self, chat_id: str, db: Optional[Session] = None) -> Optional[ChatModel]: - with get_db_context(db) as db: + async def insert_shared_chat_by_chat_id(self, chat_id: str, db: Optional[AsyncSession] = None) -> Optional[ChatModel]: + async with get_async_db_context(db) as db: # Get the existing chat to share - chat = db.get(Chat, chat_id) + chat = await db.get(Chat, chat_id) # Check if chat exists if not chat: return None # Check if the chat is already shared if chat.share_id: - return self.get_chat_by_id_and_user_id(chat.share_id, 'shared', db=db) + return await self.get_chat_by_id_and_user_id(chat.share_id, 'shared', db=db) # Create a new chat with the same data, but with a new ID shared_chat = ChatModel( **{ @@ -581,22 +577,23 @@ class ChatTable: ) shared_result = Chat(**shared_chat.model_dump()) db.add(shared_result) - db.commit() - db.refresh(shared_result) + await db.commit() + await db.refresh(shared_result) # Update the original chat with the share_id - result = db.query(Chat).filter_by(id=chat_id).update({'share_id': shared_chat.id}) - db.commit() - return shared_chat if (shared_result and result) else None + await db.execute(update(Chat).filter_by(id=chat_id).values(share_id=shared_chat.id)) + await db.commit() + return shared_chat if shared_result else None - def update_shared_chat_by_chat_id(self, chat_id: str, db: Optional[Session] = None) -> Optional[ChatModel]: + async def update_shared_chat_by_chat_id(self, chat_id: str, db: Optional[AsyncSession] = None) -> Optional[ChatModel]: try: - with get_db_context(db) as db: - chat = db.get(Chat, chat_id) - shared_chat = db.query(Chat).filter_by(user_id=f'shared-{chat_id}').first() + async with get_async_db_context(db) as db: + chat = await db.get(Chat, chat_id) + result = await db.execute(select(Chat).filter_by(user_id=f'shared-{chat_id}')) + shared_chat = result.scalars().first() if shared_chat is None: - return self.insert_shared_chat_by_chat_id(chat_id, db=db) + return await self.insert_shared_chat_by_chat_id(chat_id, db=db) shared_chat.title = chat.title shared_chat.chat = chat.chat @@ -604,99 +601,100 @@ class ChatTable: shared_chat.pinned = chat.pinned shared_chat.folder_id = chat.folder_id shared_chat.updated_at = int(time.time()) - db.commit() - db.refresh(shared_chat) + await db.commit() + await db.refresh(shared_chat) return ChatModel.model_validate(shared_chat) except Exception: return None - def delete_shared_chat_by_chat_id(self, chat_id: str, db: Optional[Session] = None) -> bool: + async def delete_shared_chat_by_chat_id(self, chat_id: str, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: - # Use subquery to delete chat_messages for shared chats - shared_chat_id_subquery = db.query(Chat.id).filter_by(user_id=f'shared-{chat_id}').scalar_subquery() - db.query(ChatMessage).filter(ChatMessage.chat_id.in_(shared_chat_id_subquery)).delete( - synchronize_session=False - ) - db.query(Chat).filter_by(user_id=f'shared-{chat_id}').delete() - db.commit() + async with get_async_db_context(db) as db: + # Get shared chat IDs + result = await db.execute(select(Chat.id).filter_by(user_id=f'shared-{chat_id}')) + shared_ids = [row[0] for row in result.all()] + + if shared_ids: + await db.execute(delete(ChatMessage).filter(ChatMessage.chat_id.in_(shared_ids))) + await db.execute(delete(Chat).filter_by(user_id=f'shared-{chat_id}')) + await db.commit() return True except Exception: return False - def unarchive_all_chats_by_user_id(self, user_id: str, db: Optional[Session] = None) -> bool: + async def unarchive_all_chats_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: - db.query(Chat).filter_by(user_id=user_id).update({'archived': False}) - db.commit() + async with get_async_db_context(db) as db: + await db.execute(update(Chat).filter_by(user_id=user_id).values(archived=False)) + await db.commit() return True except Exception: return False - def update_chat_share_id_by_id( - self, id: str, share_id: Optional[str], db: Optional[Session] = None + async def update_chat_share_id_by_id( + self, id: str, share_id: Optional[str], db: Optional[AsyncSession] = None ) -> Optional[ChatModel]: try: - with get_db_context(db) as db: - chat = db.get(Chat, id) + async with get_async_db_context(db) as db: + chat = await db.get(Chat, id) chat.share_id = share_id - db.commit() - db.refresh(chat) + await db.commit() + await db.refresh(chat) return ChatModel.model_validate(chat) except Exception: return None - def toggle_chat_pinned_by_id(self, id: str, db: Optional[Session] = None) -> Optional[ChatModel]: + async def toggle_chat_pinned_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[ChatModel]: try: - with get_db_context(db) as db: - chat = db.get(Chat, id) + async with get_async_db_context(db) as db: + chat = await db.get(Chat, id) chat.pinned = not chat.pinned chat.updated_at = int(time.time()) - db.commit() - db.refresh(chat) + await db.commit() + await db.refresh(chat) return ChatModel.model_validate(chat) except Exception: return None - def toggle_chat_archive_by_id(self, id: str, db: Optional[Session] = None) -> Optional[ChatModel]: + async def toggle_chat_archive_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[ChatModel]: try: - with get_db_context(db) as db: - chat = db.get(Chat, id) + async with get_async_db_context(db) as db: + chat = await db.get(Chat, id) chat.archived = not chat.archived chat.folder_id = None chat.updated_at = int(time.time()) - db.commit() - db.refresh(chat) + await db.commit() + await db.refresh(chat) return ChatModel.model_validate(chat) except Exception: return None - def archive_all_chats_by_user_id(self, user_id: str, db: Optional[Session] = None) -> bool: + async def archive_all_chats_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: - db.query(Chat).filter_by(user_id=user_id).update({'archived': True}) - db.commit() + async with get_async_db_context(db) as db: + await db.execute(update(Chat).filter_by(user_id=user_id).values(archived=True)) + await db.commit() return True except Exception: return False - def get_archived_chat_list_by_user_id( + async def get_archived_chat_list_by_user_id( self, user_id: str, filter: Optional[dict] = None, skip: int = 0, limit: int = 50, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[ChatTitleIdResponse]: - with get_db_context(db) as db: - query = db.query(Chat).filter_by(user_id=user_id, archived=True) + async with get_async_db_context(db) as db: + stmt = select(Chat.id, Chat.title, Chat.updated_at, Chat.created_at).filter_by(user_id=user_id, archived=True) if filter: query_key = filter.get('query') if query_key: - query = query.filter(Chat.title.ilike(f'%{query_key}%')) + stmt = stmt.filter(Chat.title.ilike(f'%{query_key}%')) order_by = filter.get('order_by') direction = filter.get('direction') @@ -706,22 +704,21 @@ class ChatTable: raise ValueError('Invalid order_by field') if direction.lower() == 'asc': - query = query.order_by(getattr(Chat, order_by).asc(), Chat.id) + stmt = stmt.order_by(getattr(Chat, order_by).asc(), Chat.id) elif direction.lower() == 'desc': - query = query.order_by(getattr(Chat, order_by).desc(), Chat.id) + stmt = stmt.order_by(getattr(Chat, order_by).desc(), Chat.id) else: raise ValueError('Invalid direction for ordering') else: - query = query.order_by(Chat.updated_at.desc(), Chat.id) - - query = query.with_entities(Chat.id, Chat.title, Chat.updated_at, Chat.created_at) + stmt = stmt.order_by(Chat.updated_at.desc(), Chat.id) if skip: - query = query.offset(skip) + stmt = stmt.offset(skip) if limit: - query = query.limit(limit) + stmt = stmt.limit(limit) - all_chats = query.all() + result = await db.execute(stmt) + all_chats = result.all() return [ ChatTitleIdResponse.model_validate( { @@ -734,21 +731,21 @@ class ChatTable: for chat in all_chats ] - def get_shared_chat_list_by_user_id( + async def get_shared_chat_list_by_user_id( self, user_id: str, filter: Optional[dict] = None, skip: int = 0, limit: int = 50, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[SharedChatResponse]: - with get_db_context(db) as db: - query = db.query(Chat).filter_by(user_id=user_id).filter(Chat.share_id.isnot(None)) + async with get_async_db_context(db) as db: + stmt = select(Chat.id, Chat.title, Chat.share_id, Chat.updated_at, Chat.created_at).filter_by(user_id=user_id).filter(Chat.share_id.isnot(None)) if filter: query_key = filter.get('query') if query_key: - query = query.filter(Chat.title.ilike(f'%{query_key}%')) + stmt = stmt.filter(Chat.title.ilike(f'%{query_key}%')) order_by = filter.get('order_by') direction = filter.get('direction') @@ -758,30 +755,21 @@ class ChatTable: raise ValueError('Invalid order_by field') if direction.lower() == 'asc': - query = query.order_by(getattr(Chat, order_by).asc(), Chat.id) + stmt = stmt.order_by(getattr(Chat, order_by).asc(), Chat.id) elif direction.lower() == 'desc': - query = query.order_by(getattr(Chat, order_by).desc(), Chat.id) + stmt = stmt.order_by(getattr(Chat, order_by).desc(), Chat.id) else: raise ValueError('Invalid direction for ordering') else: - query = query.order_by(Chat.updated_at.desc(), Chat.id) - - # Select only the columns needed for SharedChatResponse - # to avoid loading the heavy chat JSON blob - query = query.with_entities( - Chat.id, - Chat.title, - Chat.share_id, - Chat.updated_at, - Chat.created_at, - ) + stmt = stmt.order_by(Chat.updated_at.desc(), Chat.id) if skip: - query = query.offset(skip) + stmt = stmt.offset(skip) if limit: - query = query.limit(limit) + stmt = stmt.limit(limit) - all_chats = query.all() + result = await db.execute(stmt) + all_chats = result.all() return [ SharedChatResponse.model_validate( { @@ -795,46 +783,45 @@ class ChatTable: for chat in all_chats ] - def get_chat_list_by_user_id( + async def get_chat_list_by_user_id( self, user_id: str, include_archived: bool = False, filter: Optional[dict] = None, skip: int = 0, limit: int = 50, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[ChatTitleIdResponse]: - with get_db_context(db) as db: - query = db.query(Chat).filter_by(user_id=user_id) + async with get_async_db_context(db) as db: + stmt = select(Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at).filter_by(user_id=user_id) if not include_archived: - query = query.filter_by(archived=False) + stmt = stmt.filter_by(archived=False) if filter: query_key = filter.get('query') if query_key: - query = query.filter(Chat.title.ilike(f'%{query_key}%')) + stmt = stmt.filter(Chat.title.ilike(f'%{query_key}%')) order_by = filter.get('order_by') direction = filter.get('direction') if order_by and direction and getattr(Chat, order_by): if direction.lower() == 'asc': - query = query.order_by(getattr(Chat, order_by).asc(), Chat.id) + stmt = stmt.order_by(getattr(Chat, order_by).asc(), Chat.id) elif direction.lower() == 'desc': - query = query.order_by(getattr(Chat, order_by).desc(), Chat.id) + stmt = stmt.order_by(getattr(Chat, order_by).desc(), Chat.id) else: raise ValueError('Invalid direction for ordering') else: - query = query.order_by(Chat.updated_at.desc(), Chat.id) - - query = query.with_entities(Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at) + stmt = stmt.order_by(Chat.updated_at.desc(), Chat.id) if skip: - query = query.offset(skip) + stmt = stmt.offset(skip) if limit: - query = query.limit(limit) + stmt = stmt.limit(limit) - all_chats = query.all() + result = await db.execute(stmt) + all_chats = result.all() return [ ChatTitleIdResponse.model_validate( { @@ -848,7 +835,7 @@ class ChatTable: for chat in all_chats ] - def get_chat_title_id_list_by_user_id( + async def get_chat_title_id_list_by_user_id( self, user_id: str, include_archived: bool = False, @@ -856,32 +843,30 @@ class ChatTable: include_pinned: bool = False, skip: Optional[int] = None, limit: Optional[int] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[ChatTitleIdResponse]: - with get_db_context(db) as db: - query = db.query(Chat).filter_by(user_id=user_id) + async with get_async_db_context(db) as db: + stmt = select(Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at).filter_by(user_id=user_id) if not include_folders: - query = query.filter_by(folder_id=None) + stmt = stmt.filter_by(folder_id=None) if not include_pinned: - query = query.filter(or_(Chat.pinned == False, Chat.pinned == None)) + stmt = stmt.filter(or_(Chat.pinned == False, Chat.pinned == None)) if not include_archived: - query = query.filter_by(archived=False) + stmt = stmt.filter_by(archived=False) - query = query.order_by(Chat.updated_at.desc(), Chat.id).with_entities( - Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at - ) + stmt = stmt.order_by(Chat.updated_at.desc(), Chat.id) if skip: - query = query.offset(skip) + stmt = stmt.offset(skip) if limit: - query = query.limit(limit) + stmt = stmt.limit(limit) - all_chats = query.all() + result = await db.execute(stmt) + all_chats = result.all() - # result has to be destructured from sqlalchemy `row` and mapped to a dict since the `ChatModel`is not the returned dataclass. return [ ChatTitleIdResponse.model_validate( { @@ -895,106 +880,107 @@ class ChatTable: for chat in all_chats ] - def get_chat_list_by_chat_ids( + async def get_chat_list_by_chat_ids( self, chat_ids: list[str], skip: int = 0, limit: int = 50, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[ChatModel]: - with get_db_context(db) as db: - all_chats = ( - db.query(Chat) + async with get_async_db_context(db) as db: + result = await db.execute( + select(Chat) .filter(Chat.id.in_(chat_ids)) .filter_by(archived=False) .order_by(Chat.updated_at.desc()) - .all() ) + all_chats = result.scalars().all() return [ChatModel.model_validate(chat) for chat in all_chats] - def get_chat_by_id(self, id: str, db: Optional[Session] = None) -> Optional[ChatModel]: + async def get_chat_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[ChatModel]: try: - with get_db_context(db) as db: - chat_item = db.get(Chat, id) + async with get_async_db_context(db) as db: + chat_item = await db.get(Chat, id) if chat_item is None: return None if self._sanitize_chat_row(chat_item): - db.commit() - db.refresh(chat_item) + await db.commit() + await db.refresh(chat_item) return ChatModel.model_validate(chat_item) except Exception: return None - def get_chat_by_share_id(self, id: str, db: Optional[Session] = None) -> Optional[ChatModel]: + async def get_chat_by_share_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[ChatModel]: try: - with get_db_context(db) as db: - # it is possible that the shared link was deleted. hence, - # we check if the chat is still shared by checking if a chat with the share_id exists - chat = db.query(Chat).filter_by(share_id=id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(Chat).filter_by(share_id=id)) + chat = result.scalars().first() if chat: - return self.get_chat_by_id(id, db=db) + return await self.get_chat_by_id(id, db=db) else: return None except Exception: return None - def get_chat_by_id_and_user_id(self, id: str, user_id: str, db: Optional[Session] = None) -> Optional[ChatModel]: + async def get_chat_by_id_and_user_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> Optional[ChatModel]: try: - with get_db_context(db) as db: - chat = db.query(Chat).filter_by(id=id, user_id=user_id).first() - return ChatModel.model_validate(chat) + async with get_async_db_context(db) as db: + result = await db.execute(select(Chat).filter_by(id=id, user_id=user_id)) + chat = result.scalars().first() + return ChatModel.model_validate(chat) if chat else None except Exception: return None - def is_chat_owner(self, id: str, user_id: str, db: Optional[Session] = None) -> bool: + async def is_chat_owner(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> bool: """ Lightweight ownership check — uses EXISTS subquery instead of loading the full Chat row (which includes the potentially large JSON blob). """ try: - with get_db_context(db) as db: - return db.query(exists().where(and_(Chat.id == id, Chat.user_id == user_id))).scalar() + async with get_async_db_context(db) as db: + result = await db.execute( + select(exists().where(and_(Chat.id == id, Chat.user_id == user_id))) + ) + return result.scalar() except Exception: return False - def get_chat_folder_id(self, id: str, user_id: str, db: Optional[Session] = None) -> Optional[str]: + async def get_chat_folder_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> Optional[str]: """ Fetch only the folder_id column for a chat, without loading the full JSON blob. Returns None if chat doesn't exist or doesn't belong to user. """ try: - with get_db_context(db) as db: - result = db.query(Chat.folder_id).filter_by(id=id, user_id=user_id).first() - return result[0] if result else None + async with get_async_db_context(db) as db: + result = await db.execute(select(Chat.folder_id).filter_by(id=id, user_id=user_id)) + row = result.first() + return row[0] if row else None except Exception: return None - def get_chats(self, skip: int = 0, limit: int = 50, db: Optional[Session] = None) -> list[ChatModel]: - with get_db_context(db) as db: - all_chats = ( - db.query(Chat) - # .limit(limit).offset(skip) - .order_by(Chat.updated_at.desc()) - ) + async def get_chats(self, skip: int = 0, limit: int = 50, db: Optional[AsyncSession] = None) -> list[ChatModel]: + async with get_async_db_context(db) as db: + result = await db.execute(select(Chat).order_by(Chat.updated_at.desc())) + all_chats = result.scalars().all() return [ChatModel.model_validate(chat) for chat in all_chats] - def get_chats_by_user_id( + async def get_chats_by_user_id( self, user_id: str, filter: Optional[dict] = None, skip: Optional[int] = None, limit: Optional[int] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> ChatListResponse: - with get_db_context(db) as db: - query = db.query(Chat).filter_by(user_id=user_id) + async with get_async_db_context(db) as db: + stmt = select(Chat).filter_by(user_id=user_id) if filter: if filter.get('updated_at'): - query = query.filter(Chat.updated_at > filter.get('updated_at')) + stmt = stmt.filter(Chat.updated_at > filter.get('updated_at')) order_by = filter.get('order_by') direction = filter.get('direction') @@ -1002,23 +988,27 @@ class ChatTable: if order_by and direction: if hasattr(Chat, order_by): if direction.lower() == 'asc': - query = query.order_by(getattr(Chat, order_by).asc(), Chat.id) + stmt = stmt.order_by(getattr(Chat, order_by).asc(), Chat.id) elif direction.lower() == 'desc': - query = query.order_by(getattr(Chat, order_by).desc(), Chat.id) + stmt = stmt.order_by(getattr(Chat, order_by).desc(), Chat.id) else: - query = query.order_by(Chat.updated_at.desc(), Chat.id) + stmt = stmt.order_by(Chat.updated_at.desc(), Chat.id) else: - query = query.order_by(Chat.updated_at.desc(), Chat.id) + stmt = stmt.order_by(Chat.updated_at.desc(), Chat.id) - total = query.count() + count_result = await db.execute( + select(func.count()).select_from(stmt.subquery()) + ) + total = count_result.scalar() if skip is not None: - query = query.offset(skip) + stmt = stmt.offset(skip) if limit is not None: - query = query.limit(limit) + stmt = stmt.limit(limit) - all_chats = query.all() + result = await db.execute(stmt) + all_chats = result.scalars().all() return ChatListResponse( **{ @@ -1027,14 +1017,14 @@ class ChatTable: } ) - def get_pinned_chats_by_user_id(self, user_id: str, db: Optional[Session] = None) -> list[ChatTitleIdResponse]: - with get_db_context(db) as db: - all_chats = ( - db.query(Chat) + async def get_pinned_chats_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[ChatTitleIdResponse]: + async with get_async_db_context(db) as db: + result = await db.execute( + select(Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at) .filter_by(user_id=user_id, pinned=True, archived=False) .order_by(Chat.updated_at.desc()) - .with_entities(Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at) ) + all_chats = result.all() return [ ChatTitleIdResponse.model_validate( { @@ -1048,19 +1038,21 @@ class ChatTable: for chat in all_chats ] - def get_archived_chats_by_user_id(self, user_id: str, db: Optional[Session] = None) -> list[ChatModel]: - with get_db_context(db) as db: - all_chats = db.query(Chat).filter_by(user_id=user_id, archived=True).order_by(Chat.updated_at.desc()) - return [ChatModel.model_validate(chat) for chat in all_chats] + async def get_archived_chats_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[ChatModel]: + async with get_async_db_context(db) as db: + result = await db.execute( + select(Chat).filter_by(user_id=user_id, archived=True).order_by(Chat.updated_at.desc()) + ) + return [ChatModel.model_validate(chat) for chat in result.scalars().all()] - def get_chats_by_user_id_and_search_text( + async def get_chats_by_user_id_and_search_text( self, user_id: str, search_text: str, include_archived: bool = False, skip: int = 0, limit: int = 60, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[ChatModel]: """ Filters chats based on a search query using Python, allowing pagination using skip and limit. @@ -1068,17 +1060,17 @@ class ChatTable: search_text = sanitize_text_for_db(search_text).lower().strip() if not search_text: - return self.get_chat_list_by_user_id(user_id, include_archived, filter={}, skip=skip, limit=limit, db=db) + return await self.get_chat_list_by_user_id(user_id, include_archived, filter={}, skip=skip, limit=limit, db=db) search_text_words = search_text.split(' ') - # search_text might contain 'tag:tag_name' format so we need to extract the tag_name, split the search_text and remove the tags + # search_text might contain 'tag:tag_name' format so we need to extract the tag_name tag_ids = [ word.replace('tag:', '').replace(' ', '_').lower() for word in search_text_words if word.startswith('tag:') ] - # Extract folder names - handle spaces and case insensitivity - folders = Folders.search_folders_by_names( + # Extract folder names + folders = await Folders.search_folders_by_names( user_id, [word.replace('folder:', '') for word in search_text_words if word.startswith('folder:')], ) @@ -1116,30 +1108,31 @@ class ChatTable: search_text = ' '.join(search_text_words) - with get_db_context(db) as db: - query = db.query(Chat).filter(Chat.user_id == user_id) + async with get_async_db_context(db) as db: + stmt = select(Chat).filter(Chat.user_id == user_id) if is_archived is not None: - query = query.filter(Chat.archived == is_archived) + stmt = stmt.filter(Chat.archived == is_archived) elif not include_archived: - query = query.filter(Chat.archived == False) + stmt = stmt.filter(Chat.archived == False) if is_pinned is not None: - query = query.filter(Chat.pinned == is_pinned) + stmt = stmt.filter(Chat.pinned == is_pinned) if is_shared is not None: if is_shared: - query = query.filter(Chat.share_id.isnot(None)) + stmt = stmt.filter(Chat.share_id.isnot(None)) else: - query = query.filter(Chat.share_id.is_(None)) + stmt = stmt.filter(Chat.share_id.is_(None)) if folder_ids: - query = query.filter(Chat.folder_id.in_(folder_ids)) + stmt = stmt.filter(Chat.folder_id.in_(folder_ids)) - query = query.order_by(Chat.updated_at.desc(), Chat.id) + stmt = stmt.order_by(Chat.updated_at.desc(), Chat.id) # Check if the database dialect is either 'sqlite' or 'postgresql' - dialect_name = db.bind.dialect.name + bind = await db.connection() + dialect_name = bind.dialect.name if dialect_name == 'sqlite': # SQLite case: using JSON1 extension for JSON searching sqlite_content_sql = ( @@ -1150,15 +1143,15 @@ class ChatTable: ')' ) sqlite_content_clause = text(sqlite_content_sql) - query = query.filter( + stmt = stmt.filter( or_(Chat.title.ilike(bindparam('title_key')), sqlite_content_clause).params( title_key=f'%{search_text}%', content_key=search_text ) ) - # Check if there are any tags to filter, it should have all the tags + # Check if there are any tags to filter if 'none' in tag_ids: - query = query.filter( + stmt = stmt.filter( text(""" NOT EXISTS ( SELECT 1 @@ -1167,7 +1160,7 @@ class ChatTable: """) ) elif tag_ids: - query = query.filter( + stmt = stmt.filter( and_( *[ text(f""" @@ -1183,14 +1176,11 @@ class ChatTable: ) elif dialect_name == 'postgresql': - # PostgreSQL doesn't allow null bytes in text. We filter those out by checking - # the JSON representation for \u0000 before attempting text extraction - # Safety filter: JSON field must not contain \u0000 - query = query.filter(text("Chat.chat::text NOT LIKE '%\\\\u0000%'")) + stmt = stmt.filter(text("Chat.chat::text NOT LIKE '%\\\\u0000%'")) # Safety filter: title must not contain actual null bytes - query = query.filter(text("Chat.title::text NOT LIKE '%\\x00%'")) + stmt = stmt.filter(text("Chat.title::text NOT LIKE '%\\x00%'")) postgres_content_sql = """ EXISTS ( @@ -1203,16 +1193,15 @@ class ChatTable: postgres_content_clause = text(postgres_content_sql) - query = query.filter( + stmt = stmt.filter( or_( Chat.title.ilike(bindparam('title_key')), postgres_content_clause, ) ).params(title_key=f'%{search_text}%', content_key=search_text.lower()) - # Check if there are any tags to filter, it should have all the tags if 'none' in tag_ids: - query = query.filter( + stmt = stmt.filter( text(""" NOT EXISTS ( SELECT 1 @@ -1221,7 +1210,7 @@ class ChatTable: """) ) elif tag_ids: - query = query.filter( + stmt = stmt.filter( and_( *[ text(f""" @@ -1236,39 +1225,42 @@ class ChatTable: ) ) else: - raise NotImplementedError(f'Unsupported dialect: {db.bind.dialect.name}') + raise NotImplementedError(f'Unsupported dialect: {dialect_name}') # Perform pagination at the SQL level - all_chats = query.offset(skip).limit(limit).all() + stmt = stmt.offset(skip).limit(limit) + result = await db.execute(stmt) + all_chats = result.scalars().all() log.info(f'The number of chats: {len(all_chats)}') # Validate and return chats return [ChatModel.model_validate(chat) for chat in all_chats] - def get_chats_by_folder_id_and_user_id( + async def get_chats_by_folder_id_and_user_id( self, folder_id: str, user_id: str, skip: int = 0, limit: int = 60, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[ChatTitleIdResponse]: - with get_db_context(db) as db: - query = db.query(Chat).filter_by(folder_id=folder_id, user_id=user_id) - query = query.filter(or_(Chat.pinned == False, Chat.pinned == None)) - query = query.filter_by(archived=False) - - query = query.order_by(Chat.updated_at.desc(), Chat.id) - - query = query.with_entities(Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at) + async with get_async_db_context(db) as db: + stmt = ( + select(Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at) + .filter_by(folder_id=folder_id, user_id=user_id) + .filter(or_(Chat.pinned == False, Chat.pinned == None)) + .filter_by(archived=False) + .order_by(Chat.updated_at.desc(), Chat.id) + ) if skip: - query = query.offset(skip) + stmt = stmt.offset(skip) if limit: - query = query.limit(limit) + stmt = stmt.limit(limit) - all_chats = query.all() + result = await db.execute(stmt) + all_chats = result.all() return [ ChatTitleIdResponse.model_validate( { @@ -1282,76 +1274,78 @@ class ChatTable: for chat in all_chats ] - def get_chats_by_folder_ids_and_user_id( - self, folder_ids: list[str], user_id: str, db: Optional[Session] = None + async def get_chats_by_folder_ids_and_user_id( + self, folder_ids: list[str], user_id: str, db: Optional[AsyncSession] = None ) -> list[ChatModel]: - with get_db_context(db) as db: - query = db.query(Chat).filter(Chat.folder_id.in_(folder_ids), Chat.user_id == user_id) - query = query.filter(or_(Chat.pinned == False, Chat.pinned == None)) - query = query.filter_by(archived=False) + async with get_async_db_context(db) as db: + stmt = ( + select(Chat) + .filter(Chat.folder_id.in_(folder_ids), Chat.user_id == user_id) + .filter(or_(Chat.pinned == False, Chat.pinned == None)) + .filter_by(archived=False) + .order_by(Chat.updated_at.desc()) + ) - query = query.order_by(Chat.updated_at.desc()) - - all_chats = query.all() + result = await db.execute(stmt) + all_chats = result.scalars().all() return [ChatModel.model_validate(chat) for chat in all_chats] - def update_chat_folder_id_by_id_and_user_id( - self, id: str, user_id: str, folder_id: str, db: Optional[Session] = None + async def update_chat_folder_id_by_id_and_user_id( + self, id: str, user_id: str, folder_id: str, db: Optional[AsyncSession] = None ) -> Optional[ChatModel]: try: - with get_db_context(db) as db: - chat = db.get(Chat, id) + async with get_async_db_context(db) as db: + chat = await db.get(Chat, id) chat.folder_id = folder_id chat.updated_at = int(time.time()) chat.pinned = False - db.commit() - db.refresh(chat) + await db.commit() + await db.refresh(chat) return ChatModel.model_validate(chat) except Exception: return None - def get_chat_tags_by_id_and_user_id(self, id: str, user_id: str, db: Optional[Session] = None) -> list[TagModel]: - with get_db_context(db) as db: - chat = db.get(Chat, id) + async def get_chat_tags_by_id_and_user_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> list[TagModel]: + async with get_async_db_context(db) as db: + chat = await db.get(Chat, id) tag_ids = chat.meta.get('tags', []) - return Tags.get_tags_by_ids_and_user_id(tag_ids, user_id, db=db) + return await Tags.get_tags_by_ids_and_user_id(tag_ids, user_id, db=db) - def get_chat_list_by_user_id_and_tag_name( + async def get_chat_list_by_user_id_and_tag_name( self, user_id: str, tag_name: str, skip: int = 0, limit: int = 50, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[ChatTitleIdResponse]: - with get_db_context(db) as db: - query = db.query(Chat).filter_by(user_id=user_id) + async with get_async_db_context(db) as db: + stmt = select(Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at).filter_by(user_id=user_id) tag_id = tag_name.replace(' ', '_').lower() - log.info(f'DB dialect name: {db.bind.dialect.name}') - if db.bind.dialect.name == 'sqlite': - # SQLite JSON1 querying for tags within the meta JSON field - query = query.filter( + bind = await db.connection() + dialect_name = bind.dialect.name + log.info(f'DB dialect name: {dialect_name}') + if dialect_name == 'sqlite': + stmt = stmt.filter( text(f"EXISTS (SELECT 1 FROM json_each(Chat.meta, '$.tags') WHERE json_each.value = :tag_id)") ).params(tag_id=tag_id) - elif db.bind.dialect.name == 'postgresql': - # PostgreSQL JSON query for tags within the meta JSON field (for `json` type) - query = query.filter( + elif dialect_name == 'postgresql': + stmt = stmt.filter( text("EXISTS (SELECT 1 FROM json_array_elements_text(Chat.meta->'tags') elem WHERE elem = :tag_id)") ).params(tag_id=tag_id) else: - raise NotImplementedError(f'Unsupported dialect: {db.bind.dialect.name}') + raise NotImplementedError(f'Unsupported dialect: {dialect_name}') - query = query.order_by(Chat.updated_at.desc(), Chat.id) - - query = query.with_entities(Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at) + stmt = stmt.order_by(Chat.updated_at.desc(), Chat.id) if skip: - query = query.offset(skip) + stmt = stmt.offset(skip) if limit: - query = query.limit(limit) + stmt = stmt.limit(limit) - all_chats = query.all() + result = await db.execute(stmt) + all_chats = result.all() return [ ChatTitleIdResponse.model_validate( { @@ -1365,49 +1359,52 @@ class ChatTable: for chat in all_chats ] - def add_chat_tag_by_id_and_user_id_and_tag_name( - self, id: str, user_id: str, tag_name: str, db: Optional[Session] = None + async def add_chat_tag_by_id_and_user_id_and_tag_name( + self, id: str, user_id: str, tag_name: str, db: Optional[AsyncSession] = None ) -> Optional[ChatModel]: tag_id = tag_name.replace(' ', '_').lower() - Tags.ensure_tags_exist([tag_name], user_id, db=db) + await Tags.ensure_tags_exist([tag_name], user_id, db=db) try: - with get_db_context(db) as db: - chat = db.get(Chat, id) + async with get_async_db_context(db) as db: + chat = await db.get(Chat, id) if tag_id not in chat.meta.get('tags', []): chat.meta = { **chat.meta, 'tags': list(set(chat.meta.get('tags', []) + [tag_id])), } - db.commit() - db.refresh(chat) + await db.commit() + await db.refresh(chat) return ChatModel.model_validate(chat) except Exception: return None - def count_chats_by_tag_name_and_user_id(self, tag_name: str, user_id: str, db: Optional[Session] = None) -> int: - with get_db_context(db) as db: - query = db.query(Chat).filter_by(user_id=user_id, archived=False) + async def count_chats_by_tag_name_and_user_id(self, tag_name: str, user_id: str, db: Optional[AsyncSession] = None) -> int: + async with get_async_db_context(db) as db: + stmt = select(func.count(Chat.id)).filter_by(user_id=user_id, archived=False) tag_id = tag_name.replace(' ', '_').lower() - if db.bind.dialect.name == 'sqlite': - query = query.filter( + bind = await db.connection() + dialect_name = bind.dialect.name + if dialect_name == 'sqlite': + stmt = stmt.filter( text("EXISTS (SELECT 1 FROM json_each(Chat.meta, '$.tags') WHERE json_each.value = :tag_id)") ).params(tag_id=tag_id) - elif db.bind.dialect.name == 'postgresql': - query = query.filter( + elif dialect_name == 'postgresql': + stmt = stmt.filter( text("EXISTS (SELECT 1 FROM json_array_elements_text(Chat.meta->'tags') elem WHERE elem = :tag_id)") ).params(tag_id=tag_id) else: - raise NotImplementedError(f'Unsupported dialect: {db.bind.dialect.name}') + raise NotImplementedError(f'Unsupported dialect: {dialect_name}') - return query.count() + result = await db.execute(stmt) + return result.scalar() - def delete_orphan_tags_for_user( + async def delete_orphan_tags_for_user( self, tag_ids: list[str], user_id: str, threshold: int = 0, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> None: """Delete tag rows from *tag_ids* that appear in at most *threshold* non-archived chats for *user_id*. One query to find orphans, one to @@ -1419,30 +1416,30 @@ class ChatTable: """ if not tag_ids: return - with get_db_context(db) as db: + async with get_async_db_context(db) as db: orphans = [] for tag_id in tag_ids: - count = self.count_chats_by_tag_name_and_user_id(tag_id, user_id, db=db) + count = await self.count_chats_by_tag_name_and_user_id(tag_id, user_id, db=db) if count <= threshold: orphans.append(tag_id) - Tags.delete_tags_by_ids_and_user_id(orphans, user_id, db=db) + await Tags.delete_tags_by_ids_and_user_id(orphans, user_id, db=db) - def count_chats_by_folder_id_and_user_id(self, folder_id: str, user_id: str, db: Optional[Session] = None) -> int: - with get_db_context(db) as db: - query = db.query(Chat).filter_by(user_id=user_id) - - query = query.filter_by(folder_id=folder_id) - count = query.count() + async def count_chats_by_folder_id_and_user_id(self, folder_id: str, user_id: str, db: Optional[AsyncSession] = None) -> int: + async with get_async_db_context(db) as db: + result = await db.execute( + select(func.count(Chat.id)).filter_by(user_id=user_id, folder_id=folder_id) + ) + count = result.scalar() log.info(f"Count of chats for folder '{folder_id}': {count}") return count - def delete_tag_by_id_and_user_id_and_tag_name( - self, id: str, user_id: str, tag_name: str, db: Optional[Session] = None + async def delete_tag_by_id_and_user_id_and_tag_name( + self, id: str, user_id: str, tag_name: str, db: Optional[AsyncSession] = None ) -> bool: try: - with get_db_context(db) as db: - chat = db.get(Chat, id) + async with get_async_db_context(db) as db: + chat = await db.get(Chat, id) tags = chat.meta.get('tags', []) tag_id = tag_name.replace(' ', '_').lower() @@ -1451,134 +1448,140 @@ class ChatTable: **chat.meta, 'tags': list(set(tags)), } - db.commit() + await db.commit() return True except Exception: return False - def delete_all_tags_by_id_and_user_id(self, id: str, user_id: str, db: Optional[Session] = None) -> bool: + async def delete_all_tags_by_id_and_user_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: - chat = db.get(Chat, id) + async with get_async_db_context(db) as db: + chat = await db.get(Chat, id) chat.meta = { **chat.meta, 'tags': [], } - db.commit() + await db.commit() return True except Exception: return False - def delete_chat_by_id(self, id: str, db: Optional[Session] = None) -> bool: + async def delete_chat_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: - db.query(AutomationRun).filter_by(chat_id=id).update( - {AutomationRun.chat_id: None}, synchronize_session=False + async with get_async_db_context(db) as db: + await db.execute( + update(AutomationRun).filter_by(chat_id=id).values(chat_id=None) ) - db.query(ChatMessage).filter_by(chat_id=id).delete() - db.query(Chat).filter_by(id=id).delete() - db.commit() + await db.execute(delete(ChatMessage).filter_by(chat_id=id)) + await db.execute(delete(Chat).filter_by(id=id)) + await db.commit() - return True and self.delete_shared_chat_by_chat_id(id, db=db) + return True and await self.delete_shared_chat_by_chat_id(id, db=db) except Exception: return False - def delete_chat_by_id_and_user_id(self, id: str, user_id: str, db: Optional[Session] = None) -> bool: + async def delete_chat_by_id_and_user_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: - db.query(AutomationRun).filter_by(chat_id=id).update( - {AutomationRun.chat_id: None}, synchronize_session=False + async with get_async_db_context(db) as db: + await db.execute( + update(AutomationRun).filter_by(chat_id=id).values(chat_id=None) ) - db.query(ChatMessage).filter_by(chat_id=id).delete() - db.query(Chat).filter_by(id=id, user_id=user_id).delete() - db.commit() + await db.execute(delete(ChatMessage).filter_by(chat_id=id)) + await db.execute(delete(Chat).filter_by(id=id, user_id=user_id)) + await db.commit() - return True and self.delete_shared_chat_by_chat_id(id, db=db) + return True and await self.delete_shared_chat_by_chat_id(id, db=db) except Exception: return False - def delete_chats_by_user_id(self, user_id: str, db: Optional[Session] = None) -> bool: + async def delete_chats_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: - self.delete_shared_chats_by_user_id(user_id, db=db) + async with get_async_db_context(db) as db: + await self.delete_shared_chats_by_user_id(user_id, db=db) - chat_id_subquery = db.query(Chat.id).filter_by(user_id=user_id).subquery() - db.query(AutomationRun).filter(AutomationRun.chat_id.in_(chat_id_subquery)).update( - {AutomationRun.chat_id: None}, synchronize_session=False + chat_id_subquery = select(Chat.id).filter_by(user_id=user_id).scalar_subquery() + await db.execute( + update(AutomationRun).filter(AutomationRun.chat_id.in_(select(Chat.id).filter_by(user_id=user_id))).values(chat_id=None) ) - db.query(ChatMessage).filter(ChatMessage.chat_id.in_(chat_id_subquery)).delete( - synchronize_session=False + await db.execute( + delete(ChatMessage).filter(ChatMessage.chat_id.in_(select(Chat.id).filter_by(user_id=user_id))) ) - db.query(Chat).filter_by(user_id=user_id).delete() - db.commit() + await db.execute(delete(Chat).filter_by(user_id=user_id)) + await db.commit() return True except Exception: return False - def delete_chats_by_user_id_and_folder_id(self, user_id: str, folder_id: str, db: Optional[Session] = None) -> bool: + async def delete_chats_by_user_id_and_folder_id(self, user_id: str, folder_id: str, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: - chat_id_subquery = db.query(Chat.id).filter_by(user_id=user_id, folder_id=folder_id).subquery() - db.query(AutomationRun).filter(AutomationRun.chat_id.in_(chat_id_subquery)).update( - {AutomationRun.chat_id: None}, synchronize_session=False + async with get_async_db_context(db) as db: + chat_ids_stmt = select(Chat.id).filter_by(user_id=user_id, folder_id=folder_id) + await db.execute( + update(AutomationRun).filter(AutomationRun.chat_id.in_(chat_ids_stmt)).values(chat_id=None) ) - db.query(ChatMessage).filter(ChatMessage.chat_id.in_(chat_id_subquery)).delete( - synchronize_session=False + await db.execute( + delete(ChatMessage).filter(ChatMessage.chat_id.in_(chat_ids_stmt)) ) - db.query(Chat).filter_by(user_id=user_id, folder_id=folder_id).delete() - db.commit() + await db.execute(delete(Chat).filter_by(user_id=user_id, folder_id=folder_id)) + await db.commit() return True except Exception: return False - def move_chats_by_user_id_and_folder_id( + async def move_chats_by_user_id_and_folder_id( self, user_id: str, folder_id: str, new_folder_id: Optional[str], - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> bool: try: - with get_db_context(db) as db: - db.query(Chat).filter_by(user_id=user_id, folder_id=folder_id).update({'folder_id': new_folder_id}) - db.commit() + async with get_async_db_context(db) as db: + await db.execute( + update(Chat).filter_by(user_id=user_id, folder_id=folder_id).values(folder_id=new_folder_id) + ) + await db.commit() return True except Exception: return False - def delete_shared_chats_by_user_id(self, user_id: str, db: Optional[Session] = None) -> bool: + async def delete_shared_chats_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: - id_rows = db.query(Chat.id).filter_by(user_id=user_id).all() + async with get_async_db_context(db) as db: + result = await db.execute(select(Chat.id).filter_by(user_id=user_id)) + id_rows = result.all() shared_chat_ids = [f'shared-{row[0]}' for row in id_rows] - # Use subquery to delete chat_messages for shared chats - shared_id_subq = db.query(Chat.id).filter(Chat.user_id.in_(shared_chat_ids)).subquery() - db.query(ChatMessage).filter(ChatMessage.chat_id.in_(shared_id_subq)).delete(synchronize_session=False) - db.query(Chat).filter(Chat.user_id.in_(shared_chat_ids)).delete() - db.commit() + if shared_chat_ids: + # Get shared chat IDs to delete associated messages + shared_result = await db.execute(select(Chat.id).filter(Chat.user_id.in_(shared_chat_ids))) + shared_ids = [row[0] for row in shared_result.all()] + if shared_ids: + await db.execute(delete(ChatMessage).filter(ChatMessage.chat_id.in_(shared_ids))) + await db.execute(delete(Chat).filter(Chat.user_id.in_(shared_chat_ids))) + await db.commit() return True except Exception: return False - def insert_chat_files( + async def insert_chat_files( self, chat_id: str, message_id: str, file_ids: list[str], user_id: str, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[list[ChatFileModel]]: if not file_ids: return None chat_message_file_ids = [ - item.id for item in self.get_chat_files_by_chat_id_and_message_id(chat_id, message_id, db=db) + item.id for item in await self.get_chat_files_by_chat_id_and_message_id(chat_id, message_id, db=db) ] # Remove duplicates and existing file_ids file_ids = list(set([file_id for file_id in file_ids if file_id and file_id not in chat_message_file_ids])) @@ -1586,7 +1589,7 @@ class ChatTable: return None try: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: now = int(time.time()) chat_files = [ @@ -1605,66 +1608,66 @@ class ChatTable: results = [ChatFile(**chat_file.model_dump()) for chat_file in chat_files] db.add_all(results) - db.commit() + await db.commit() return chat_files except Exception: return None - def get_chat_files_by_chat_id_and_message_id( - self, chat_id: str, message_id: str, db: Optional[Session] = None + async def get_chat_files_by_chat_id_and_message_id( + self, chat_id: str, message_id: str, db: Optional[AsyncSession] = None ) -> list[ChatFileModel]: - with get_db_context(db) as db: - all_chat_files = ( - db.query(ChatFile) + async with get_async_db_context(db) as db: + result = await db.execute( + select(ChatFile) .filter_by(chat_id=chat_id, message_id=message_id) .order_by(ChatFile.created_at.asc()) - .all() ) + all_chat_files = result.scalars().all() return [ChatFileModel.model_validate(chat_file) for chat_file in all_chat_files] - def delete_chat_file(self, chat_id: str, file_id: str, db: Optional[Session] = None) -> bool: + async def delete_chat_file(self, chat_id: str, file_id: str, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: - db.query(ChatFile).filter_by(chat_id=chat_id, file_id=file_id).delete() - db.commit() + async with get_async_db_context(db) as db: + await db.execute(delete(ChatFile).filter_by(chat_id=chat_id, file_id=file_id)) + await db.commit() return True except Exception: return False - def get_shared_chats_by_file_id(self, file_id: str, db: Optional[Session] = None) -> list[ChatModel]: - with get_db_context(db) as db: - # Join Chat and ChatFile tables to get shared chats associated with the file_id - all_chats = ( - db.query(Chat) + async def get_shared_chats_by_file_id(self, file_id: str, db: Optional[AsyncSession] = None) -> list[ChatModel]: + async with get_async_db_context(db) as db: + result = await db.execute( + select(Chat) .join(ChatFile, Chat.id == ChatFile.chat_id) .filter(ChatFile.file_id == file_id, Chat.share_id.isnot(None)) - .all() ) + all_chats = result.scalars().all() return [ChatModel.model_validate(chat) for chat in all_chats] - def update_chat_tasks_by_id(self, id: str, tasks: list[dict]) -> Optional[ChatModel]: + async def update_chat_tasks_by_id(self, id: str, tasks: list[dict]) -> Optional[ChatModel]: """Update the tasks list on a chat.""" try: - with get_db_context() as db: - chat = db.get(Chat, id) + async with get_async_db_context() as db: + chat = await db.get(Chat, id) if chat is None: return None chat.tasks = tasks - db.commit() - db.refresh(chat) + await db.commit() + await db.refresh(chat) return ChatModel.model_validate(chat) except Exception: return None - def get_chat_tasks_by_id(self, id: str) -> list[dict]: + async def get_chat_tasks_by_id(self, id: str) -> list[dict]: """Read the tasks list from a chat (lightweight column query).""" - with get_db_context() as db: - result = db.query(Chat.tasks).filter_by(id=id).first() - if result is None or result[0] is None: + async with get_async_db_context() as db: + result = await db.execute(select(Chat.tasks).filter_by(id=id)) + row = result.first() + if row is None or row[0] is None: return [] - return result[0] + return row[0] Chats = ChatTable() diff --git a/backend/open_webui/models/feedbacks.py b/backend/open_webui/models/feedbacks.py index 9172e2ba8e..61124619b5 100644 --- a/backend/open_webui/models/feedbacks.py +++ b/backend/open_webui/models/feedbacks.py @@ -3,9 +3,10 @@ import time import uuid from typing import Optional -from sqlalchemy.orm import Session -from open_webui.internal.db import Base, JSONField, get_db, get_db_context -from open_webui.models.users import User +from sqlalchemy import select, delete, func +from sqlalchemy.ext.asyncio import AsyncSession +from open_webui.internal.db import Base, JSONField, get_async_db_context +from open_webui.models.users import User, UserModel from pydantic import BaseModel, ConfigDict from sqlalchemy import BigInteger, Column, Text, JSON, Boolean @@ -139,10 +140,10 @@ class ModelHistoryResponse(BaseModel): class FeedbackTable: - def insert_new_feedback( - self, user_id: str, form_data: FeedbackForm, db: Optional[Session] = None + async def insert_new_feedback( + self, user_id: str, form_data: FeedbackForm, db: Optional[AsyncSession] = None ) -> Optional[FeedbackModel]: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: id = str(uuid.uuid4()) feedback = FeedbackModel( **{ @@ -157,8 +158,8 @@ class FeedbackTable: try: result = Feedback(**feedback.model_dump()) db.add(result) - db.commit() - db.refresh(result) + await db.commit() + await db.refresh(result) if result: return FeedbackModel.model_validate(result) else: @@ -167,97 +168,101 @@ class FeedbackTable: log.exception(f'Error creating a new feedback: {e}') return None - def get_feedback_by_id(self, id: str, db: Optional[Session] = None) -> Optional[FeedbackModel]: + async def get_feedback_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[FeedbackModel]: try: - with get_db_context(db) as db: - feedback = db.query(Feedback).filter_by(id=id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(Feedback).filter_by(id=id)) + feedback = result.scalars().first() if not feedback: return None return FeedbackModel.model_validate(feedback) except Exception: return None - def get_feedback_by_id_and_user_id( - self, id: str, user_id: str, db: Optional[Session] = None + async def get_feedback_by_id_and_user_id( + self, id: str, user_id: str, db: Optional[AsyncSession] = None ) -> Optional[FeedbackModel]: try: - with get_db_context(db) as db: - feedback = db.query(Feedback).filter_by(id=id, user_id=user_id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(Feedback).filter_by(id=id, user_id=user_id)) + feedback = result.scalars().first() if not feedback: return None return FeedbackModel.model_validate(feedback) except Exception: return None - def get_feedbacks_by_chat_id(self, chat_id: str, db: Optional[Session] = None) -> list[FeedbackModel]: + async def get_feedbacks_by_chat_id(self, chat_id: str, db: Optional[AsyncSession] = None) -> list[FeedbackModel]: """Get all feedbacks for a specific chat.""" try: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: # meta.chat_id stores the chat reference - feedbacks = ( - db.query(Feedback) + result = await db.execute( + select(Feedback) .filter(Feedback.meta['chat_id'].as_string() == chat_id) .order_by(Feedback.created_at.desc()) - .all() ) + feedbacks = result.scalars().all() return [FeedbackModel.model_validate(fb) for fb in feedbacks] except Exception: return [] - def get_feedback_items( + async def get_feedback_items( self, filter: dict = {}, skip: int = 0, limit: int = 30, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> FeedbackListResponse: - with get_db_context(db) as db: - query = db.query(Feedback, User).join(User, Feedback.user_id == User.id) + async with get_async_db_context(db) as db: + stmt = select(Feedback, User).join(User, Feedback.user_id == User.id) if filter: # Apply model_id filter (exact match) model_id = filter.get('model_id') if model_id: - query = query.filter(Feedback.data['model_id'].as_string() == model_id) + stmt = stmt.filter(Feedback.data['model_id'].as_string() == model_id) order_by = filter.get('order_by') direction = filter.get('direction') if order_by == 'username': if direction == 'asc': - query = query.order_by(User.name.asc()) + stmt = stmt.order_by(User.name.asc()) else: - query = query.order_by(User.name.desc()) + stmt = stmt.order_by(User.name.desc()) elif order_by == 'model_id': - # it's stored in feedback.data['model_id'] if direction == 'asc': - query = query.order_by(Feedback.data['model_id'].as_string().asc()) + stmt = stmt.order_by(Feedback.data['model_id'].as_string().asc()) else: - query = query.order_by(Feedback.data['model_id'].as_string().desc()) + stmt = stmt.order_by(Feedback.data['model_id'].as_string().desc()) elif order_by == 'rating': - # it's stored in feedback.data['rating'] if direction == 'asc': - query = query.order_by(Feedback.data['rating'].as_string().asc()) + stmt = stmt.order_by(Feedback.data['rating'].as_string().asc()) else: - query = query.order_by(Feedback.data['rating'].as_string().desc()) + stmt = stmt.order_by(Feedback.data['rating'].as_string().desc()) elif order_by == 'updated_at': if direction == 'asc': - query = query.order_by(Feedback.updated_at.asc()) + stmt = stmt.order_by(Feedback.updated_at.asc()) else: - query = query.order_by(Feedback.updated_at.desc()) + stmt = stmt.order_by(Feedback.updated_at.desc()) else: - query = query.order_by(Feedback.created_at.desc()) + stmt = stmt.order_by(Feedback.created_at.desc()) # Count BEFORE pagination - total = query.count() + count_result = await db.execute( + select(func.count()).select_from(stmt.subquery()) + ) + total = count_result.scalar() if skip: - query = query.offset(skip) + stmt = stmt.offset(skip) if limit: - query = query.limit(limit) + stmt = stmt.limit(limit) - items = query.all() + result = await db.execute(stmt) + items = result.all() feedbacks = [] for feedback, user in items: @@ -267,15 +272,17 @@ class FeedbackTable: return FeedbackListResponse(items=feedbacks, total=total) - def get_all_feedbacks(self, db: Optional[Session] = None) -> list[FeedbackModel]: - with get_db_context(db) as db: - return [ - FeedbackModel.model_validate(feedback) - for feedback in db.query(Feedback).order_by(Feedback.updated_at.desc()).all() - ] + async def get_all_feedbacks(self, db: Optional[AsyncSession] = None) -> list[FeedbackModel]: + async with get_async_db_context(db) as db: + result = await db.execute(select(Feedback).order_by(Feedback.updated_at.desc())) + return [FeedbackModel.model_validate(feedback) for feedback in result.scalars().all()] - def get_all_feedback_ids(self, db: Optional[Session] = None) -> list[FeedbackIdResponse]: - with get_db_context(db) as db: + async def get_all_feedback_ids(self, db: Optional[AsyncSession] = None) -> list[FeedbackIdResponse]: + async with get_async_db_context(db) as db: + result = await db.execute( + select(Feedback.id, Feedback.user_id, Feedback.created_at, Feedback.updated_at) + .order_by(Feedback.updated_at.desc()) + ) return [ FeedbackIdResponse( id=row.id, @@ -283,36 +290,28 @@ class FeedbackTable: created_at=row.created_at, updated_at=row.updated_at, ) - for row in db.query( - Feedback.id, - Feedback.user_id, - Feedback.created_at, - Feedback.updated_at, - ) - .order_by(Feedback.updated_at.desc()) - .all() + for row in result.all() ] - def get_distinct_model_ids(self, db: Optional[Session] = None) -> list[str]: + async def get_distinct_model_ids(self, db: Optional[AsyncSession] = None) -> list[str]: """Get distinct model_ids from feedback data for filter dropdowns.""" - with get_db_context(db) as db: - rows = ( - db.query(Feedback.data['model_id'].as_string()) + async with get_async_db_context(db) as db: + result = await db.execute( + select(Feedback.data['model_id'].as_string()) .filter(Feedback.data['model_id'].as_string().isnot(None)) .distinct() - .all() ) + rows = result.all() return sorted([row[0] for row in rows if row[0]]) - def get_feedbacks_for_leaderboard(self, db: Optional[Session] = None) -> list[LeaderboardFeedbackData]: + async def get_feedbacks_for_leaderboard(self, db: Optional[AsyncSession] = None) -> list[LeaderboardFeedbackData]: """Fetch only id and data for leaderboard computation (excludes snapshot/meta).""" - with get_db_context(db) as db: - return [ - LeaderboardFeedbackData(id=row.id, data=row.data) for row in db.query(Feedback.id, Feedback.data).all() - ] + async with get_async_db_context(db) as db: + result = await db.execute(select(Feedback.id, Feedback.data)) + return [LeaderboardFeedbackData(id=row.id, data=row.data) for row in result.all()] - def get_model_evaluation_history( - self, model_id: str, days: int = 30, db: Optional[Session] = None + async def get_model_evaluation_history( + self, model_id: str, days: int = 30, db: Optional[AsyncSession] = None ) -> list[ModelHistoryEntry]: """ Get daily wins/losses for a specific model over the past N days. @@ -322,13 +321,16 @@ class FeedbackTable: from datetime import datetime, timedelta from collections import defaultdict - with get_db_context(db) as db: + async with get_async_db_context(db) as db: if days == 0: # All time - no cutoff - rows = db.query(Feedback.created_at, Feedback.data).all() + result = await db.execute(select(Feedback.created_at, Feedback.data)) else: cutoff = int(time.time()) - (days * 86400) - rows = db.query(Feedback.created_at, Feedback.data).filter(Feedback.created_at >= cutoff).all() + result = await db.execute( + select(Feedback.created_at, Feedback.data).filter(Feedback.created_at >= cutoff) + ) + rows = result.all() daily_counts = defaultdict(lambda: {'won': 0, 'lost': 0}) first_date = None @@ -374,25 +376,26 @@ class FeedbackTable: return result - def get_feedbacks_by_type(self, type: str, db: Optional[Session] = None) -> list[FeedbackModel]: - with get_db_context(db) as db: - return [ - FeedbackModel.model_validate(feedback) - for feedback in db.query(Feedback).filter_by(type=type).order_by(Feedback.updated_at.desc()).all() - ] + async def get_feedbacks_by_type(self, type: str, db: Optional[AsyncSession] = None) -> list[FeedbackModel]: + async with get_async_db_context(db) as db: + result = await db.execute( + select(Feedback).filter_by(type=type).order_by(Feedback.updated_at.desc()) + ) + return [FeedbackModel.model_validate(feedback) for feedback in result.scalars().all()] - def get_feedbacks_by_user_id(self, user_id: str, db: Optional[Session] = None) -> list[FeedbackModel]: - with get_db_context(db) as db: - return [ - FeedbackModel.model_validate(feedback) - for feedback in db.query(Feedback).filter_by(user_id=user_id).order_by(Feedback.updated_at.desc()).all() - ] + async def get_feedbacks_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[FeedbackModel]: + async with get_async_db_context(db) as db: + result = await db.execute( + select(Feedback).filter_by(user_id=user_id).order_by(Feedback.updated_at.desc()) + ) + return [FeedbackModel.model_validate(feedback) for feedback in result.scalars().all()] - def update_feedback_by_id( - self, id: str, form_data: FeedbackForm, db: Optional[Session] = None + async def update_feedback_by_id( + self, id: str, form_data: FeedbackForm, db: Optional[AsyncSession] = None ) -> Optional[FeedbackModel]: - with get_db_context(db) as db: - feedback = db.query(Feedback).filter_by(id=id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(Feedback).filter_by(id=id)) + feedback = result.scalars().first() if not feedback: return None @@ -405,18 +408,19 @@ class FeedbackTable: feedback.updated_at = int(time.time()) - db.commit() + await db.commit() return FeedbackModel.model_validate(feedback) - def update_feedback_by_id_and_user_id( + async def update_feedback_by_id_and_user_id( self, id: str, user_id: str, form_data: FeedbackForm, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[FeedbackModel]: - with get_db_context(db) as db: - feedback = db.query(Feedback).filter_by(id=id, user_id=user_id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(Feedback).filter_by(id=id, user_id=user_id)) + feedback = result.scalars().first() if not feedback: return None @@ -429,38 +433,40 @@ class FeedbackTable: feedback.updated_at = int(time.time()) - db.commit() + await db.commit() return FeedbackModel.model_validate(feedback) - def delete_feedback_by_id(self, id: str, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: - feedback = db.query(Feedback).filter_by(id=id).first() + async def delete_feedback_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: + result = await db.execute(select(Feedback).filter_by(id=id)) + feedback = result.scalars().first() if not feedback: return False - db.delete(feedback) - db.commit() + await db.delete(feedback) + await db.commit() return True - def delete_feedback_by_id_and_user_id(self, id: str, user_id: str, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: - feedback = db.query(Feedback).filter_by(id=id, user_id=user_id).first() + async def delete_feedback_by_id_and_user_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: + result = await db.execute(select(Feedback).filter_by(id=id, user_id=user_id)) + feedback = result.scalars().first() if not feedback: return False - db.delete(feedback) - db.commit() + await db.delete(feedback) + await db.commit() return True - def delete_feedbacks_by_user_id(self, user_id: str, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: - result = db.query(Feedback).filter_by(user_id=user_id).delete() - db.commit() - return result > 0 + async def delete_feedbacks_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: + result = await db.execute(delete(Feedback).filter_by(user_id=user_id)) + await db.commit() + return result.rowcount > 0 - def delete_all_feedbacks(self, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: - result = db.query(Feedback).delete() - db.commit() - return result > 0 + async def delete_all_feedbacks(self, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: + result = await db.execute(delete(Feedback)) + await db.commit() + return result.rowcount > 0 Feedbacks = FeedbackTable() diff --git a/backend/open_webui/models/files.py b/backend/open_webui/models/files.py index 7a9f77a3b0..f79255f50b 100644 --- a/backend/open_webui/models/files.py +++ b/backend/open_webui/models/files.py @@ -2,8 +2,9 @@ import logging import time from typing import Optional -from sqlalchemy.orm import Session -from open_webui.internal.db import Base, JSONField, get_db, get_db_context +from sqlalchemy import select, delete, func +from sqlalchemy.ext.asyncio import AsyncSession +from open_webui.internal.db import Base, JSONField, get_async_db_context from open_webui.utils.misc import sanitize_metadata from pydantic import BaseModel, ConfigDict, model_validator from sqlalchemy import BigInteger, Column, String, Text, JSON @@ -124,8 +125,8 @@ class FileUpdateForm(BaseModel): class FilesTable: - def insert_new_file(self, user_id: str, form_data: FileForm, db: Optional[Session] = None) -> Optional[FileModel]: - with get_db_context(db) as db: + async def insert_new_file(self, user_id: str, form_data: FileForm, db: Optional[AsyncSession] = None) -> Optional[FileModel]: + async with get_async_db_context(db) as db: file_data = form_data.model_dump() # Sanitize meta to remove non-JSON-serializable objects @@ -145,8 +146,8 @@ class FilesTable: try: result = File(**file.model_dump()) db.add(result) - db.commit() - db.refresh(result) + await db.commit() + await db.refresh(result) if result: return FileModel.model_validate(result) else: @@ -155,21 +156,22 @@ class FilesTable: log.exception(f'Error inserting a new file: {e}') return None - def get_file_by_id(self, id: str, db: Optional[Session] = None) -> Optional[FileModel]: + async def get_file_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[FileModel]: try: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: try: - file = db.get(File, id) - return FileModel.model_validate(file) + file = await db.get(File, id) + return FileModel.model_validate(file) if file else None except Exception: return None except Exception: return None - def get_file_by_id_and_user_id(self, id: str, user_id: str, db: Optional[Session] = None) -> Optional[FileModel]: - with get_db_context(db) as db: + async def get_file_by_id_and_user_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> Optional[FileModel]: + async with get_async_db_context(db) as db: try: - file = db.query(File).filter_by(id=id, user_id=user_id).first() + result = await db.execute(select(File).filter_by(id=id, user_id=user_id)) + file = result.scalars().first() if file: return FileModel.model_validate(file) else: @@ -177,10 +179,12 @@ class FilesTable: except Exception: return None - def get_file_metadata_by_id(self, id: str, db: Optional[Session] = None) -> Optional[FileMetadataResponse]: - with get_db_context(db) as db: + async def get_file_metadata_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[FileMetadataResponse]: + async with get_async_db_context(db) as db: try: - file = db.get(File, id) + file = await db.get(File, id) + if not file: + return None return FileMetadataResponse( id=file.id, hash=file.hash, @@ -191,12 +195,13 @@ class FilesTable: except Exception: return None - def get_files(self, db: Optional[Session] = None) -> list[FileModel]: - with get_db_context(db) as db: - return [FileModel.model_validate(file) for file in db.query(File).all()] + async def get_files(self, db: Optional[AsyncSession] = None) -> list[FileModel]: + async with get_async_db_context(db) as db: + result = await db.execute(select(File)) + return [FileModel.model_validate(file) for file in result.scalars().all()] - def check_access_by_user_id(self, id, user_id, permission='write', db: Optional[Session] = None) -> bool: - file = self.get_file_by_id(id, db=db) + async def check_access_by_user_id(self, id, user_id, permission='write', db: Optional[AsyncSession] = None) -> bool: + file = await self.get_file_by_id(id, db=db) if not file: return False if file.user_id == user_id: @@ -204,50 +209,59 @@ class FilesTable: # Implement additional access control logic here as needed return False - def get_files_by_ids(self, ids: list[str], db: Optional[Session] = None) -> list[FileModel]: - with get_db_context(db) as db: - return [ - FileModel.model_validate(file) - for file in db.query(File).filter(File.id.in_(ids)).order_by(File.updated_at.desc()).all() - ] + async def get_files_by_ids(self, ids: list[str], db: Optional[AsyncSession] = None) -> list[FileModel]: + async with get_async_db_context(db) as db: + result = await db.execute( + select(File).filter(File.id.in_(ids)).order_by(File.updated_at.desc()) + ) + return [FileModel.model_validate(file) for file in result.scalars().all()] - def get_file_metadatas_by_ids(self, ids: list[str], db: Optional[Session] = None) -> list[FileMetadataResponse]: - with get_db_context(db) as db: - return [ - FileMetadataResponse( - id=file.id, - hash=file.hash, - meta=file.meta, - created_at=file.created_at, - updated_at=file.updated_at, - ) - for file in db.query(File.id, File.hash, File.meta, File.created_at, File.updated_at) + async def get_file_metadatas_by_ids(self, ids: list[str], db: Optional[AsyncSession] = None) -> list[FileMetadataResponse]: + async with get_async_db_context(db) as db: + result = await db.execute( + select(File.id, File.hash, File.meta, File.created_at, File.updated_at) .filter(File.id.in_(ids)) .order_by(File.updated_at.desc()) - .all() + ) + return [ + FileMetadataResponse( + id=row.id, + hash=row.hash, + meta=row.meta, + created_at=row.created_at, + updated_at=row.updated_at, + ) + for row in result.all() ] - def get_files_by_user_id(self, user_id: str, db: Optional[Session] = None) -> list[FileModel]: - with get_db_context(db) as db: - return [FileModel.model_validate(file) for file in db.query(File).filter_by(user_id=user_id).all()] + async def get_files_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[FileModel]: + async with get_async_db_context(db) as db: + result = await db.execute(select(File).filter_by(user_id=user_id)) + return [FileModel.model_validate(file) for file in result.scalars().all()] - def get_file_list( + async def get_file_list( self, user_id: Optional[str] = None, skip: int = 0, limit: int = 50, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> 'FileListResponse': - with get_db_context(db) as db: - query = db.query(File) + async with get_async_db_context(db) as db: + stmt = select(File) if user_id: - query = query.filter_by(user_id=user_id) + stmt = stmt.filter_by(user_id=user_id) - total = query.count() + count_result = await db.execute( + select(func.count()).select_from(stmt.subquery()) + ) + total = count_result.scalar() + result = await db.execute( + stmt.order_by(File.updated_at.desc(), File.id.desc()).offset(skip).limit(limit) + ) items = [ FileModelResponse.model_validate(file, from_attributes=True) - for file in query.order_by(File.updated_at.desc(), File.id.desc()).offset(skip).limit(limit).all() + for file in result.scalars().all() ] return FileListResponse(items=items, total=total) @@ -275,13 +289,13 @@ class FilesTable: pattern = pattern.replace('?', '_') return pattern - def search_files( + async def search_files( self, user_id: Optional[str] = None, filename: str = '*', skip: int = 0, limit: int = 100, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[FileModel]: """ Search files with glob pattern matching, optional user filter, and pagination. @@ -296,27 +310,28 @@ class FilesTable: Returns: List of matching FileModel objects, ordered by created_at descending. """ - with get_db_context(db) as db: - query = db.query(File) + async with get_async_db_context(db) as db: + stmt = select(File) if user_id: - query = query.filter_by(user_id=user_id) + stmt = stmt.filter_by(user_id=user_id) pattern = self._glob_to_like_pattern(filename) if pattern != '%': - query = query.filter(File.filename.ilike(pattern, escape='\\')) + stmt = stmt.filter(File.filename.ilike(pattern, escape='\\')) - return [ - FileModel.model_validate(file) - for file in query.order_by(File.created_at.desc(), File.id.desc()).offset(skip).limit(limit).all() - ] + result = await db.execute( + stmt.order_by(File.created_at.desc(), File.id.desc()).offset(skip).limit(limit) + ) + return [FileModel.model_validate(file) for file in result.scalars().all()] - def update_file_by_id( - self, id: str, form_data: FileUpdateForm, db: Optional[Session] = None + async def update_file_by_id( + self, id: str, form_data: FileUpdateForm, db: Optional[AsyncSession] = None ) -> Optional[FileModel]: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: try: - file = db.query(File).filter_by(id=id).first() + result = await db.execute(select(File).filter_by(id=id)) + file = result.scalars().first() if form_data.hash is not None: file.hash = form_data.hash @@ -328,63 +343,64 @@ class FilesTable: file.meta = {**(file.meta if file.meta else {}), **form_data.meta} file.updated_at = int(time.time()) - db.commit() + await db.commit() return FileModel.model_validate(file) except Exception as e: log.exception(f'Error updating file completely by id: {e}') return None - def update_file_hash_by_id(self, id: str, hash: Optional[str], db: Optional[Session] = None) -> Optional[FileModel]: - with get_db_context(db) as db: + async def update_file_hash_by_id(self, id: str, hash: Optional[str], db: Optional[AsyncSession] = None) -> Optional[FileModel]: + async with get_async_db_context(db) as db: try: - file = db.query(File).filter_by(id=id).first() + result = await db.execute(select(File).filter_by(id=id)) + file = result.scalars().first() file.hash = hash file.updated_at = int(time.time()) - db.commit() + await db.commit() return FileModel.model_validate(file) except Exception: return None - def update_file_data_by_id(self, id: str, data: dict, db: Optional[Session] = None) -> Optional[FileModel]: - with get_db_context(db) as db: + async def update_file_data_by_id(self, id: str, data: dict, db: Optional[AsyncSession] = None) -> Optional[FileModel]: + async with get_async_db_context(db) as db: try: - file = db.query(File).filter_by(id=id).first() + result = await db.execute(select(File).filter_by(id=id)) + file = result.scalars().first() file.data = {**(file.data if file.data else {}), **data} file.updated_at = int(time.time()) - db.commit() + await db.commit() return FileModel.model_validate(file) except Exception as e: return None - def update_file_metadata_by_id(self, id: str, meta: dict, db: Optional[Session] = None) -> Optional[FileModel]: - with get_db_context(db) as db: + async def update_file_metadata_by_id(self, id: str, meta: dict, db: Optional[AsyncSession] = None) -> Optional[FileModel]: + async with get_async_db_context(db) as db: try: - file = db.query(File).filter_by(id=id).first() + result = await db.execute(select(File).filter_by(id=id)) + file = result.scalars().first() file.meta = {**(file.meta if file.meta else {}), **meta} file.updated_at = int(time.time()) - db.commit() + await db.commit() return FileModel.model_validate(file) except Exception: return None - return False - - def delete_file_by_id(self, id: str, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: + async def delete_file_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: try: - db.query(File).filter_by(id=id).delete() - db.commit() + await db.execute(delete(File).filter_by(id=id)) + await db.commit() return True except Exception: return False - def delete_all_files(self, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: + async def delete_all_files(self, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: try: - db.query(File).delete() - db.commit() + await db.execute(delete(File)) + await db.commit() return True except Exception: diff --git a/backend/open_webui/models/folders.py b/backend/open_webui/models/folders.py index cd9c9bbc67..4e2a4e9f38 100644 --- a/backend/open_webui/models/folders.py +++ b/backend/open_webui/models/folders.py @@ -6,10 +6,10 @@ import re from pydantic import BaseModel, ConfigDict -from sqlalchemy import BigInteger, Column, Text, JSON, Boolean, func -from sqlalchemy.orm import Session +from sqlalchemy import BigInteger, Column, Text, JSON, Boolean, func, select, delete +from sqlalchemy.ext.asyncio import AsyncSession -from open_webui.internal.db import Base, JSONField, get_db, get_db_context +from open_webui.internal.db import Base, JSONField, get_async_db_context log = logging.getLogger(__name__) @@ -85,14 +85,14 @@ class FolderUpdateForm(BaseModel): class FolderTable: - def insert_new_folder( + async def insert_new_folder( self, user_id: str, form_data: FolderForm, parent_id: Optional[str] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[FolderModel]: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: id = str(uuid.uuid4()) folder = FolderModel( **{ @@ -107,8 +107,8 @@ class FolderTable: try: result = Folder(**folder.model_dump()) db.add(result) - db.commit() - db.refresh(result) + await db.commit() + await db.refresh(result) if result: return FolderModel.model_validate(result) else: @@ -117,12 +117,13 @@ class FolderTable: log.exception(f'Error inserting a new folder: {e}') return None - def get_folder_by_id_and_user_id( - self, id: str, user_id: str, db: Optional[Session] = None + async def get_folder_by_id_and_user_id( + self, id: str, user_id: str, db: Optional[AsyncSession] = None ) -> Optional[FolderModel]: try: - with get_db_context(db) as db: - folder = db.query(Folder).filter_by(id=id, user_id=user_id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(Folder).filter_by(id=id, user_id=user_id)) + folder = result.scalars().first() if not folder: return None @@ -131,48 +132,50 @@ class FolderTable: except Exception: return None - def get_children_folders_by_id_and_user_id( - self, id: str, user_id: str, db: Optional[Session] = None + async def get_children_folders_by_id_and_user_id( + self, id: str, user_id: str, db: Optional[AsyncSession] = None ) -> Optional[list[FolderModel]]: try: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: folders = [] - def get_children(folder): - children = self.get_folders_by_parent_id_and_user_id(folder.id, user_id, db=db) + async def get_children(folder): + children = await self.get_folders_by_parent_id_and_user_id(folder.id, user_id, db=db) for child in children: - get_children(child) + await get_children(child) folders.append(child) - folder = db.query(Folder).filter_by(id=id, user_id=user_id).first() + result = await db.execute(select(Folder).filter_by(id=id, user_id=user_id)) + folder = result.scalars().first() if not folder: return None - get_children(folder) + await get_children(folder) return folders except Exception: return None - def get_folders_by_user_id(self, user_id: str, db: Optional[Session] = None) -> list[FolderModel]: - with get_db_context(db) as db: - return [FolderModel.model_validate(folder) for folder in db.query(Folder).filter_by(user_id=user_id).all()] + async def get_folders_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[FolderModel]: + async with get_async_db_context(db) as db: + result = await db.execute(select(Folder).filter_by(user_id=user_id)) + return [FolderModel.model_validate(folder) for folder in result.scalars().all()] - def get_folder_by_parent_id_and_user_id_and_name( + async def get_folder_by_parent_id_and_user_id_and_name( self, parent_id: Optional[str], user_id: str, name: str, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[FolderModel]: try: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: # Check if folder exists - folder = ( - db.query(Folder) + result = await db.execute( + select(Folder) .filter_by(parent_id=parent_id, user_id=user_id) .filter(Folder.name.ilike(name)) - .first() ) + folder = result.scalars().first() if not folder: return None @@ -182,25 +185,24 @@ class FolderTable: log.error(f'get_folder_by_parent_id_and_user_id_and_name: {e}') return None - def get_folders_by_parent_id_and_user_id( - self, parent_id: Optional[str], user_id: str, db: Optional[Session] = None + async def get_folders_by_parent_id_and_user_id( + self, parent_id: Optional[str], user_id: str, db: Optional[AsyncSession] = None ) -> list[FolderModel]: - with get_db_context(db) as db: - return [ - FolderModel.model_validate(folder) - for folder in db.query(Folder).filter_by(parent_id=parent_id, user_id=user_id).all() - ] + async with get_async_db_context(db) as db: + result = await db.execute(select(Folder).filter_by(parent_id=parent_id, user_id=user_id)) + return [FolderModel.model_validate(folder) for folder in result.scalars().all()] - def update_folder_parent_id_by_id_and_user_id( + async def update_folder_parent_id_by_id_and_user_id( self, id: str, user_id: str, parent_id: str, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[FolderModel]: try: - with get_db_context(db) as db: - folder = db.query(Folder).filter_by(id=id, user_id=user_id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(Folder).filter_by(id=id, user_id=user_id)) + folder = result.scalars().first() if not folder: return None @@ -208,38 +210,39 @@ class FolderTable: folder.parent_id = parent_id folder.updated_at = int(time.time()) - db.commit() + await db.commit() return FolderModel.model_validate(folder) except Exception as e: log.error(f'update_folder: {e}') return - def update_folder_by_id_and_user_id( + async def update_folder_by_id_and_user_id( self, id: str, user_id: str, form_data: FolderUpdateForm, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[FolderModel]: try: - with get_db_context(db) as db: - folder = db.query(Folder).filter_by(id=id, user_id=user_id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(Folder).filter_by(id=id, user_id=user_id)) + folder = result.scalars().first() if not folder: return None form_data = form_data.model_dump(exclude_unset=True) - existing_folder = ( - db.query(Folder) + existing_result = await db.execute( + select(Folder) .filter_by( name=form_data.get('name'), parent_id=folder.parent_id, user_id=user_id, ) - .first() ) + existing_folder = existing_result.scalars().first() if existing_folder and existing_folder.id != id: return None @@ -258,19 +261,20 @@ class FolderTable: } folder.updated_at = int(time.time()) - db.commit() + await db.commit() return FolderModel.model_validate(folder) except Exception as e: log.error(f'update_folder: {e}') return - def update_folder_is_expanded_by_id_and_user_id( - self, id: str, user_id: str, is_expanded: bool, db: Optional[Session] = None + async def update_folder_is_expanded_by_id_and_user_id( + self, id: str, user_id: str, is_expanded: bool, db: Optional[AsyncSession] = None ) -> Optional[FolderModel]: try: - with get_db_context(db) as db: - folder = db.query(Folder).filter_by(id=id, user_id=user_id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(Folder).filter_by(id=id, user_id=user_id)) + folder = result.scalars().first() if not folder: return None @@ -278,37 +282,39 @@ class FolderTable: folder.is_expanded = is_expanded folder.updated_at = int(time.time()) - db.commit() + await db.commit() return FolderModel.model_validate(folder) except Exception as e: log.error(f'update_folder: {e}') return - def delete_folder_by_id_and_user_id(self, id: str, user_id: str, db: Optional[Session] = None) -> list[str]: + async def delete_folder_by_id_and_user_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> list[str]: try: folder_ids = [] - with get_db_context(db) as db: - folder = db.query(Folder).filter_by(id=id, user_id=user_id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(Folder).filter_by(id=id, user_id=user_id)) + folder = result.scalars().first() if not folder: return folder_ids folder_ids.append(folder.id) # Delete all children folders - def delete_children(folder): - folder_children = self.get_folders_by_parent_id_and_user_id(folder.id, user_id, db=db) + async def delete_children(folder): + folder_children = await self.get_folders_by_parent_id_and_user_id(folder.id, user_id, db=db) for folder_child in folder_children: - delete_children(folder_child) + await delete_children(folder_child) folder_ids.append(folder_child.id) - folder = db.query(Folder).filter_by(id=folder_child.id).first() - db.delete(folder) - db.commit() + child_result = await db.execute(select(Folder).filter_by(id=folder_child.id)) + child_folder = child_result.scalars().first() + await db.delete(child_folder) + await db.commit() - delete_children(folder) - db.delete(folder) - db.commit() + await delete_children(folder) + await db.delete(folder) + await db.commit() return folder_ids except Exception as e: log.error(f'delete_folder: {e}') @@ -319,8 +325,8 @@ class FolderTable: name = re.sub(r'[\s_]+', ' ', name) return name.strip().lower() - def search_folders_by_names( - self, user_id: str, queries: list[str], db: Optional[Session] = None + async def search_folders_by_names( + self, user_id: str, queries: list[str], db: Optional[AsyncSession] = None ) -> list[FolderModel]: """ Search for folders for a user where the name matches any of the queries, treating _ and space as equivalent, case-insensitive. @@ -330,16 +336,18 @@ class FolderTable: return [] results = {} - with get_db_context(db) as db: - folders = db.query(Folder).filter_by(user_id=user_id).all() + async with get_async_db_context(db) as db: + result = await db.execute(select(Folder).filter_by(user_id=user_id)) + folders = result.scalars().all() for folder in folders: if self.normalize_folder_name(folder.name) in normalized_queries: results[folder.id] = FolderModel.model_validate(folder) # get children folders - children = self.get_children_folders_by_id_and_user_id(folder.id, user_id, db=db) - for child in children: - results[child.id] = child + children = await self.get_children_folders_by_id_and_user_id(folder.id, user_id, db=db) + if children: + for child in children: + results[child.id] = child # Return the results as a list if not results: @@ -348,16 +356,17 @@ class FolderTable: results = list(results.values()) return results - def search_folders_by_name_contains( - self, user_id: str, query: str, db: Optional[Session] = None + async def search_folders_by_name_contains( + self, user_id: str, query: str, db: Optional[AsyncSession] = None ) -> list[FolderModel]: """ Partial match: normalized name contains (as substring) the normalized query. """ normalized_query = self.normalize_folder_name(query) results = [] - with get_db_context(db) as db: - folders = db.query(Folder).filter_by(user_id=user_id).all() + async with get_async_db_context(db) as db: + result = await db.execute(select(Folder).filter_by(user_id=user_id)) + folders = result.scalars().all() for folder in folders: norm_name = self.normalize_folder_name(folder.name) if normalized_query in norm_name: diff --git a/backend/open_webui/models/functions.py b/backend/open_webui/models/functions.py index f9761e947a..db34454b43 100644 --- a/backend/open_webui/models/functions.py +++ b/backend/open_webui/models/functions.py @@ -2,8 +2,9 @@ import logging import time from typing import Optional -from sqlalchemy.orm import Session, defer -from open_webui.internal.db import Base, JSONField, get_db, get_db_context +from sqlalchemy import select, delete, update +from sqlalchemy.ext.asyncio import AsyncSession +from open_webui.internal.db import Base, JSONField, get_async_db_context from open_webui.models.users import Users, UserModel, UserResponse from pydantic import BaseModel, ConfigDict from sqlalchemy import BigInteger, Boolean, Column, String, Text, Index @@ -107,12 +108,12 @@ class FunctionValves(BaseModel): class FunctionsTable: - def insert_new_function( + async def insert_new_function( self, user_id: str, type: str, form_data: FunctionForm, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[FunctionModel]: function = FunctionModel( **{ @@ -125,11 +126,11 @@ class FunctionsTable: ) try: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: result = Function(**function.model_dump()) db.add(result) - db.commit() - db.refresh(result) + await db.commit() + await db.refresh(result) if result: return FunctionModel.model_validate(result) else: @@ -138,17 +139,18 @@ class FunctionsTable: log.exception(f'Error creating a new function: {e}') return None - def sync_functions( + async def sync_functions( self, user_id: str, functions: list[FunctionWithValvesModel], - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[FunctionWithValvesModel]: # Synchronize functions for a user by updating existing ones, inserting new ones, and removing those that are no longer present. try: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: # Get existing functions - existing_functions = db.query(Function).all() + result = await db.execute(select(Function)) + existing_functions = result.scalars().all() existing_ids = {func.id for func in existing_functions} # Prepare a set of new function IDs @@ -157,12 +159,12 @@ class FunctionsTable: # Update or insert functions for func in functions: if func.id in existing_ids: - db.query(Function).filter_by(id=func.id).update( - { + await db.execute( + update(Function).filter_by(id=func.id).values( **func.model_dump(), - 'user_id': user_id, - 'updated_at': int(time.time()), - } + user_id=user_id, + updated_at=int(time.time()), + ) ) else: new_func = Function( @@ -177,24 +179,25 @@ class FunctionsTable: # Remove functions that are no longer present for func in existing_functions: if func.id not in new_function_ids: - db.delete(func) + await db.delete(func) - db.commit() + await db.commit() - return [FunctionModel.model_validate(func) for func in db.query(Function).all()] + result = await db.execute(select(Function)) + return [FunctionModel.model_validate(func) for func in result.scalars().all()] except Exception as e: log.exception(f'Error syncing functions for user {user_id}: {e}') return [] - def get_function_by_id(self, id: str, db: Optional[Session] = None) -> Optional[FunctionModel]: + async def get_function_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[FunctionModel]: try: - with get_db_context(db) as db: - function = db.get(Function, id) - return FunctionModel.model_validate(function) + async with get_async_db_context(db) as db: + function = await db.get(Function, id) + return FunctionModel.model_validate(function) if function else None except Exception: return None - def get_functions_by_ids(self, ids: list[str], db: Optional[Session] = None) -> list[FunctionModel]: + async def get_functions_by_ids(self, ids: list[str], db: Optional[AsyncSession] = None) -> list[FunctionModel]: """ Batch fetch multiple functions by their IDs in a single query. Returns functions in the same order as the input IDs (None entries filtered out). @@ -202,8 +205,9 @@ class FunctionsTable: if not ids: return [] try: - with get_db_context(db) as db: - functions = db.query(Function).filter(Function.id.in_(ids)).all() + async with get_async_db_context(db) as db: + result = await db.execute(select(Function).filter(Function.id.in_(ids))) + functions = result.scalars().all() # Create a dict for O(1) lookup func_dict = {f.id: FunctionModel.model_validate(f) for f in functions} # Return in original order, filtering out any not found @@ -211,27 +215,31 @@ class FunctionsTable: except Exception: return [] - def get_functions( - self, active_only=False, include_valves=False, db: Optional[Session] = None + async def get_functions( + self, active_only=False, include_valves=False, db: Optional[AsyncSession] = None ) -> list[FunctionModel | FunctionWithValvesModel]: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: if active_only: - functions = db.query(Function).filter_by(is_active=True).all() - + result = await db.execute(select(Function).filter_by(is_active=True)) else: - functions = db.query(Function).all() + result = await db.execute(select(Function)) + + functions = result.scalars().all() if include_valves: return [FunctionWithValvesModel.model_validate(function) for function in functions] else: return [FunctionModel.model_validate(function) for function in functions] - def get_function_list(self, db: Optional[Session] = None) -> list[FunctionUserResponse]: - with get_db_context(db) as db: - functions = db.query(Function).options(defer(Function.content)).order_by(Function.updated_at.desc()).all() + async def get_function_list(self, db: Optional[AsyncSession] = None) -> list[FunctionUserResponse]: + async with get_async_db_context(db) as db: + result = await db.execute( + select(Function).order_by(Function.updated_at.desc()) + ) + functions = result.scalars().all() user_ids = list(set(func.user_id for func in functions)) - users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else [] + users = await Users.get_users_by_user_ids(user_ids, db=db) if user_ids else [] users_dict = {user.id: user for user in users} return [ @@ -253,42 +261,34 @@ class FunctionsTable: for func in functions ] - def get_functions_by_type(self, type: str, active_only=False, db: Optional[Session] = None) -> list[FunctionModel]: - with get_db_context(db) as db: + async def get_functions_by_type(self, type: str, active_only=False, db: Optional[AsyncSession] = None) -> list[FunctionModel]: + async with get_async_db_context(db) as db: if active_only: - return [ - FunctionModel.model_validate(function) - for function in db.query(Function).filter_by(type=type, is_active=True).all() - ] + result = await db.execute(select(Function).filter_by(type=type, is_active=True)) else: - return [ - FunctionModel.model_validate(function) for function in db.query(Function).filter_by(type=type).all() - ] + result = await db.execute(select(Function).filter_by(type=type)) + return [FunctionModel.model_validate(function) for function in result.scalars().all()] - def get_global_filter_functions(self, db: Optional[Session] = None) -> list[FunctionModel]: - with get_db_context(db) as db: - return [ - FunctionModel.model_validate(function) - for function in db.query(Function).filter_by(type='filter', is_active=True, is_global=True).all() - ] + async def get_global_filter_functions(self, db: Optional[AsyncSession] = None) -> list[FunctionModel]: + async with get_async_db_context(db) as db: + result = await db.execute(select(Function).filter_by(type='filter', is_active=True, is_global=True)) + return [FunctionModel.model_validate(function) for function in result.scalars().all()] - def get_global_action_functions(self, db: Optional[Session] = None) -> list[FunctionModel]: - with get_db_context(db) as db: - return [ - FunctionModel.model_validate(function) - for function in db.query(Function).filter_by(type='action', is_active=True, is_global=True).all() - ] + async def get_global_action_functions(self, db: Optional[AsyncSession] = None) -> list[FunctionModel]: + async with get_async_db_context(db) as db: + result = await db.execute(select(Function).filter_by(type='action', is_active=True, is_global=True)) + return [FunctionModel.model_validate(function) for function in result.scalars().all()] - def get_function_valves_by_id(self, id: str, db: Optional[Session] = None) -> Optional[dict]: - with get_db_context(db) as db: + async def get_function_valves_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[dict]: + async with get_async_db_context(db) as db: try: - function = db.get(Function, id) + function = await db.get(Function, id) return function.valves if function.valves else {} except Exception as e: log.exception(f'Error getting function valves by id {id}: {e}') return None - def get_function_valves_by_ids(self, ids: list[str], db: Optional[Session] = None) -> dict[str, dict]: + async def get_function_valves_by_ids(self, ids: list[str], db: Optional[AsyncSession] = None) -> dict[str, dict]: """ Batch fetch valves for multiple functions in a single query. Returns a dict mapping function_id -> valves dict. @@ -297,33 +297,34 @@ class FunctionsTable: if not ids: return {} try: - with get_db_context(db) as db: - functions = db.query(Function.id, Function.valves).filter(Function.id.in_(ids)).all() + async with get_async_db_context(db) as db: + result = await db.execute(select(Function.id, Function.valves).filter(Function.id.in_(ids))) + functions = result.all() return {f.id: (f.valves if f.valves else {}) for f in functions} except Exception as e: log.exception(f'Error batch-fetching function valves: {e}') return {} - def update_function_valves_by_id( - self, id: str, valves: dict, db: Optional[Session] = None + async def update_function_valves_by_id( + self, id: str, valves: dict, db: Optional[AsyncSession] = None ) -> Optional[FunctionValves]: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: try: - function = db.get(Function, id) + function = await db.get(Function, id) function.valves = valves function.updated_at = int(time.time()) - db.commit() - db.refresh(function) + await db.commit() + await db.refresh(function) return FunctionModel.model_validate(function) except Exception: return None - def update_function_metadata_by_id( - self, id: str, metadata: dict, db: Optional[Session] = None + async def update_function_metadata_by_id( + self, id: str, metadata: dict, db: Optional[AsyncSession] = None ) -> Optional[FunctionModel]: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: try: - function = db.get(Function, id) + function = await db.get(Function, id) if function: if function.meta: @@ -332,8 +333,8 @@ class FunctionsTable: function.meta = metadata function.updated_at = int(time.time()) - db.commit() - db.refresh(function) + await db.commit() + await db.refresh(function) return FunctionModel.model_validate(function) else: return None @@ -341,9 +342,9 @@ class FunctionsTable: log.exception(f'Error updating function metadata by id {id}: {e}') return None - def get_user_valves_by_id_and_user_id(self, id: str, user_id: str, db: Optional[Session] = None) -> Optional[dict]: + async def get_user_valves_by_id_and_user_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> Optional[dict]: try: - user = Users.get_user_by_id(user_id, db=db) + user = await Users.get_user_by_id(user_id, db=db) user_settings = user.settings.model_dump() if user.settings else {} # Check if user has "functions" and "valves" settings @@ -357,11 +358,11 @@ class FunctionsTable: log.exception(f'Error getting user values by id {id} and user id {user_id}') return None - def update_user_valves_by_id_and_user_id( - self, id: str, user_id: str, valves: dict, db: Optional[Session] = None + async def update_user_valves_by_id_and_user_id( + self, id: str, user_id: str, valves: dict, db: Optional[AsyncSession] = None ) -> Optional[dict]: try: - user = Users.get_user_by_id(user_id, db=db) + user = await Users.get_user_by_id(user_id, db=db) user_settings = user.settings.model_dump() if user.settings else {} # Check if user has "functions" and "valves" settings @@ -373,47 +374,47 @@ class FunctionsTable: user_settings['functions']['valves'][id] = valves # Update the user settings in the database - Users.update_user_by_id(user_id, {'settings': user_settings}, db=db) + await Users.update_user_by_id(user_id, {'settings': user_settings}, db=db) return user_settings['functions']['valves'][id] except Exception as e: log.exception(f'Error updating user valves by id {id} and user_id {user_id}: {e}') return None - def update_function_by_id(self, id: str, updated: dict, db: Optional[Session] = None) -> Optional[FunctionModel]: - with get_db_context(db) as db: + async def update_function_by_id(self, id: str, updated: dict, db: Optional[AsyncSession] = None) -> Optional[FunctionModel]: + async with get_async_db_context(db) as db: try: - db.query(Function).filter_by(id=id).update( - { + await db.execute( + update(Function).filter_by(id=id).values( **updated, - 'updated_at': int(time.time()), - } + updated_at=int(time.time()), + ) ) - db.commit() - function = db.get(Function, id) + await db.commit() + function = await db.get(Function, id) return FunctionModel.model_validate(function) if function else None except Exception: return None - def deactivate_all_functions(self, db: Optional[Session] = None) -> Optional[bool]: - with get_db_context(db) as db: + async def deactivate_all_functions(self, db: Optional[AsyncSession] = None) -> Optional[bool]: + async with get_async_db_context(db) as db: try: - db.query(Function).update( - { - 'is_active': False, - 'updated_at': int(time.time()), - } + await db.execute( + update(Function).values( + is_active=False, + updated_at=int(time.time()), + ) ) - db.commit() + await db.commit() return True except Exception: return None - def delete_function_by_id(self, id: str, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: + async def delete_function_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: try: - db.query(Function).filter_by(id=id).delete() - db.commit() + await db.execute(delete(Function).filter_by(id=id)) + await db.commit() return True except Exception: diff --git a/backend/open_webui/models/groups.py b/backend/open_webui/models/groups.py index fc4cfb0d31..bca9908580 100644 --- a/backend/open_webui/models/groups.py +++ b/backend/open_webui/models/groups.py @@ -4,8 +4,9 @@ import time from typing import Optional import uuid -from sqlalchemy.orm import Session -from open_webui.internal.db import Base, JSONField, get_db, get_db_context +from sqlalchemy import select, delete, update, func, and_, or_, cast, String +from sqlalchemy.ext.asyncio import AsyncSession +from open_webui.internal.db import Base, JSONField, get_async_db_context from open_webui.env import DEFAULT_GROUP_SHARE_PERMISSION from open_webui.models.files import FileMetadataResponse @@ -15,15 +16,9 @@ from pydantic import BaseModel, ConfigDict from sqlalchemy import ( BigInteger, Column, - String, Text, JSON, - and_, - func, ForeignKey, - cast, - or_, - select, ) log = logging.getLogger(__name__) @@ -143,10 +138,10 @@ class GroupTable: group_data['data']['config']['share'] = DEFAULT_GROUP_SHARE_PERMISSION return group_data - def insert_new_group( - self, user_id: str, form_data: GroupForm, db: Optional[Session] = None + async def insert_new_group( + self, user_id: str, form_data: GroupForm, db: Optional[AsyncSession] = None ) -> Optional[GroupModel]: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: group_data = self._ensure_default_share_config(form_data.model_dump(exclude_none=True)) group = GroupModel( **{ @@ -161,8 +156,8 @@ class GroupTable: try: result = Group(**group.model_dump()) db.add(result) - db.commit() - db.refresh(result) + await db.commit() + await db.refresh(result) if result: return GroupModel.model_validate(result) else: @@ -171,18 +166,20 @@ class GroupTable: except Exception: return None - def get_all_groups(self, db: Optional[Session] = None) -> list[GroupModel]: - with get_db_context(db) as db: - groups = db.query(Group).order_by(Group.updated_at.desc()).all() + async def get_all_groups(self, db: Optional[AsyncSession] = None) -> list[GroupModel]: + async with get_async_db_context(db) as db: + result = await db.execute(select(Group).order_by(Group.updated_at.desc())) + groups = result.scalars().all() return [GroupModel.model_validate(group) for group in groups] - def get_group_by_name(self, name: str, db: Optional[Session] = None) -> Optional[GroupModel]: - with get_db_context(db) as db: - group = db.query(Group).filter(Group.name == name).first() + async def get_group_by_name(self, name: str, db: Optional[AsyncSession] = None) -> Optional[GroupModel]: + async with get_async_db_context(db) as db: + result = await db.execute(select(Group).filter(Group.name == name)) + group = result.scalars().first() return GroupModel.model_validate(group) if group else None - def get_groups(self, filter, db: Optional[Session] = None) -> list[GroupResponse]: - with get_db_context(db) as db: + async def get_groups(self, filter, db: Optional[AsyncSession] = None) -> list[GroupResponse]: + async with get_async_db_context(db) as db: member_count = ( select(func.count(GroupMember.user_id)) .where(GroupMember.group_id == Group.id) @@ -190,11 +187,11 @@ class GroupTable: .scalar_subquery() .label('member_count') ) - query = db.query(Group, member_count) + stmt = select(Group, member_count) if filter: if 'query' in filter: - query = query.filter(Group.name.ilike(f'%{filter["query"]}%')) + stmt = stmt.filter(Group.name.ilike(f'%{filter["query"]}%')) # When share filter is present, member check is handled in the share logic if 'share' in filter: @@ -218,20 +215,21 @@ class GroupTable: json_share_lower == 'members', Group.id.in_(member_groups_select), ) - query = query.filter(or_(anyone_can_share, members_only_and_is_member)) + stmt = stmt.filter(or_(anyone_can_share, members_only_and_is_member)) else: - query = query.filter(anyone_can_share) + stmt = stmt.filter(anyone_can_share) else: - query = query.filter(and_(Group.data.isnot(None), json_share_lower == 'false')) + stmt = stmt.filter(and_(Group.data.isnot(None), json_share_lower == 'false')) else: # Only apply member_id filter when share filter is NOT present if 'member_id' in filter: - query = query.filter( + stmt = stmt.filter( Group.id.in_(select(GroupMember.group_id).where(GroupMember.user_id == filter['member_id'])) ) - results = query.order_by(Group.updated_at.desc()).all() + result = await db.execute(stmt.order_by(Group.updated_at.desc())) + rows = result.all() return [ GroupResponse.model_validate( @@ -240,32 +238,36 @@ class GroupTable: 'member_count': count or 0, } ) - for group, count in results + for group, count in rows ] - def search_groups( + async def search_groups( self, filter: Optional[dict] = None, skip: int = 0, limit: int = 30, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> GroupListResponse: - with get_db_context(db) as db: - query = db.query(Group) + async with get_async_db_context(db) as db: + stmt = select(Group) if filter: if 'query' in filter: - query = query.filter(Group.name.ilike(f'%{filter["query"]}%')) + stmt = stmt.filter(Group.name.ilike(f'%{filter["query"]}%')) if 'member_id' in filter: - query = query.filter( + stmt = stmt.filter( Group.id.in_(select(GroupMember.group_id).where(GroupMember.user_id == filter['member_id'])) ) if 'share' in filter: share_value = filter['share'] - query = query.filter(Group.data.op('->>')('share') == str(share_value)) + stmt = stmt.filter(Group.data.op('->>') ('share') == str(share_value)) - total = query.count() + # Get total count + count_result = await db.execute( + select(func.count()).select_from(stmt.subquery()) + ) + total = count_result.scalar() member_count = ( select(func.count(GroupMember.user_id)) @@ -274,7 +276,14 @@ class GroupTable: .scalar_subquery() .label('member_count') ) - results = query.add_columns(member_count).order_by(Group.updated_at.desc()).offset(skip).limit(limit).all() + result = await db.execute( + select(Group, member_count) + .where(Group.id.in_(select(stmt.subquery().c.id))) + .order_by(Group.updated_at.desc()) + .offset(skip) + .limit(limit) + ) + rows = result.all() return { 'items': [ @@ -284,65 +293,67 @@ class GroupTable: 'member_count': count or 0, } ) - for group, count in results + for group, count in rows ], 'total': total, } - def get_groups_by_member_id(self, user_id: str, db: Optional[Session] = None) -> list[GroupModel]: - with get_db_context(db) as db: - return [ - GroupModel.model_validate(group) - for group in db.query(Group) + async def get_groups_by_member_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[GroupModel]: + async with get_async_db_context(db) as db: + result = await db.execute( + select(Group) .join(GroupMember, GroupMember.group_id == Group.id) .filter(GroupMember.user_id == user_id) .order_by(Group.updated_at.desc()) - .all() - ] + ) + return [GroupModel.model_validate(group) for group in result.scalars().all()] - def get_groups_by_member_ids( - self, user_ids: list[str], db: Optional[Session] = None + async def get_groups_by_member_ids( + self, user_ids: list[str], db: Optional[AsyncSession] = None ) -> dict[str, list[GroupModel]]: """Fetch groups for multiple users in a single query to avoid N+1.""" - with get_db_context(db) as db: + async with get_async_db_context(db) as db: # Query GroupMember joined with Group, filtering by user_ids - results = ( - db.query(GroupMember.user_id, Group) + result = await db.execute( + select(GroupMember.user_id, Group) .join(Group, Group.id == GroupMember.group_id) .filter(GroupMember.user_id.in_(user_ids)) .order_by(Group.updated_at.desc()) - .all() ) + rows = result.all() # Group groups by user_id user_groups: dict[str, list[GroupModel]] = {uid: [] for uid in user_ids} - for user_id, group in results: + for user_id, group in rows: user_groups[user_id].append(GroupModel.model_validate(group)) return user_groups - def get_group_by_id(self, id: str, db: Optional[Session] = None) -> Optional[GroupModel]: + async def get_group_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[GroupModel]: try: - with get_db_context(db) as db: - group = db.query(Group).filter_by(id=id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(Group).filter_by(id=id)) + group = result.scalars().first() return GroupModel.model_validate(group) if group else None except Exception: return None - def get_group_user_ids_by_id(self, id: str, db: Optional[Session] = None) -> list[str]: - with get_db_context(db) as db: - members = db.query(GroupMember.user_id).filter(GroupMember.group_id == id).all() + async def get_group_user_ids_by_id(self, id: str, db: Optional[AsyncSession] = None) -> list[str]: + async with get_async_db_context(db) as db: + result = await db.execute(select(GroupMember.user_id).filter(GroupMember.group_id == id)) + members = result.all() if not members: return [] return [m[0] for m in members] - def get_group_user_ids_by_ids(self, group_ids: list[str], db: Optional[Session] = None) -> dict[str, list[str]]: - with get_db_context(db) as db: - members = ( - db.query(GroupMember.group_id, GroupMember.user_id).filter(GroupMember.group_id.in_(group_ids)).all() + async def get_group_user_ids_by_ids(self, group_ids: list[str], db: Optional[AsyncSession] = None) -> dict[str, list[str]]: + async with get_async_db_context(db) as db: + result = await db.execute( + select(GroupMember.group_id, GroupMember.user_id).filter(GroupMember.group_id.in_(group_ids)) ) + members = result.all() group_user_ids: dict[str, list[str]] = {group_id: [] for group_id in group_ids} @@ -351,10 +362,10 @@ class GroupTable: return group_user_ids - def set_group_user_ids_by_id(self, group_id: str, user_ids: list[str], db: Optional[Session] = None) -> None: - with get_db_context(db) as db: + async def set_group_user_ids_by_id(self, group_id: str, user_ids: list[str], db: Optional[AsyncSession] = None) -> None: + async with get_async_db_context(db) as db: # Delete existing members - db.query(GroupMember).filter(GroupMember.group_id == group_id).delete() + await db.execute(delete(GroupMember).filter(GroupMember.group_id == group_id)) # Insert new members now = int(time.time()) @@ -370,101 +381,106 @@ class GroupTable: ] db.add_all(new_members) - db.commit() + await db.commit() - def get_group_member_count_by_id(self, id: str, db: Optional[Session] = None) -> int: - with get_db_context(db) as db: - count = db.query(func.count(GroupMember.user_id)).filter(GroupMember.group_id == id).scalar() + async def get_group_member_count_by_id(self, id: str, db: Optional[AsyncSession] = None) -> int: + async with get_async_db_context(db) as db: + result = await db.execute(select(func.count(GroupMember.user_id)).filter(GroupMember.group_id == id)) + count = result.scalar() return count if count else 0 - def get_group_member_counts_by_ids(self, ids: list[str], db: Optional[Session] = None) -> dict[str, int]: + async def get_group_member_counts_by_ids(self, ids: list[str], db: Optional[AsyncSession] = None) -> dict[str, int]: if not ids: return {} - with get_db_context(db) as db: - rows = ( - db.query(GroupMember.group_id, func.count(GroupMember.user_id)) + async with get_async_db_context(db) as db: + result = await db.execute( + select(GroupMember.group_id, func.count(GroupMember.user_id)) .filter(GroupMember.group_id.in_(ids)) .group_by(GroupMember.group_id) - .all() ) + rows = result.all() return {group_id: count for group_id, count in rows} - def update_group_by_id( + async def update_group_by_id( self, id: str, form_data: GroupUpdateForm, overwrite: bool = False, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[GroupModel]: try: - with get_db_context(db) as db: - db.query(Group).filter_by(id=id).update( - { + async with get_async_db_context(db) as db: + await db.execute( + update(Group).filter_by(id=id).values( **form_data.model_dump(exclude_none=True), - 'updated_at': int(time.time()), - } + updated_at=int(time.time()), + ) ) - db.commit() - return self.get_group_by_id(id=id, db=db) + await db.commit() + return await self.get_group_by_id(id=id, db=db) except Exception as e: log.exception(e) return None - def delete_group_by_id(self, id: str, db: Optional[Session] = None) -> bool: + async def delete_group_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: - db.query(Group).filter_by(id=id).delete() - db.commit() + async with get_async_db_context(db) as db: + await db.execute(delete(Group).filter_by(id=id)) + await db.commit() return True except Exception: return False - def delete_all_groups(self, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: + async def delete_all_groups(self, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: try: - db.query(Group).delete() - db.commit() + await db.execute(delete(Group)) + await db.commit() return True except Exception: return False - def remove_user_from_all_groups(self, user_id: str, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: + async def remove_user_from_all_groups(self, user_id: str, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: try: # Find all groups the user belongs to - groups = ( - db.query(Group) + result = await db.execute( + select(Group) .join(GroupMember, GroupMember.group_id == Group.id) .filter(GroupMember.user_id == user_id) - .all() ) + groups = result.scalars().all() # Remove the user from each group for group in groups: - db.query(GroupMember).filter( - GroupMember.group_id == group.id, GroupMember.user_id == user_id - ).delete() + await db.execute( + delete(GroupMember).filter( + GroupMember.group_id == group.id, GroupMember.user_id == user_id + ) + ) - db.query(Group).filter_by(id=group.id).update({'updated_at': int(time.time())}) + await db.execute( + update(Group).filter_by(id=group.id).values(updated_at=int(time.time())) + ) - db.commit() + await db.commit() return True except Exception: - db.rollback() + await db.rollback() return False - def create_groups_by_group_names( - self, user_id: str, group_names: list[str], db: Optional[Session] = None + async def create_groups_by_group_names( + self, user_id: str, group_names: list[str], db: Optional[AsyncSession] = None ) -> list[GroupModel]: # check for existing groups - existing_groups = self.get_all_groups(db=db) + existing_groups = await self.get_all_groups(db=db) existing_group_names = {group.name for group in existing_groups} new_groups = [] - with get_db_context(db) as db: + async with get_async_db_context(db) as db: for group_name in group_names: if group_name not in existing_group_names: new_group = GroupModel( @@ -483,31 +499,31 @@ class GroupTable: try: result = Group(**new_group.model_dump()) db.add(result) - db.commit() - db.refresh(result) + await db.commit() + await db.refresh(result) new_groups.append(GroupModel.model_validate(result)) except Exception as e: log.exception(e) continue return new_groups - def sync_groups_by_group_names(self, user_id: str, group_names: list[str], db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: + async def sync_groups_by_group_names(self, user_id: str, group_names: list[str], db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: try: now = int(time.time()) # 1. Groups that SHOULD contain the user - target_groups = db.query(Group).filter(Group.name.in_(group_names)).all() + result = await db.execute(select(Group).filter(Group.name.in_(group_names))) + target_groups = result.scalars().all() target_group_ids = {g.id for g in target_groups} # 2. Groups the user is CURRENTLY in - existing_group_ids = { - g.id - for g in db.query(Group) + result = await db.execute( + select(Group) .join(GroupMember, GroupMember.group_id == Group.id) .filter(GroupMember.user_id == user_id) - .all() - } + ) + existing_group_ids = {g.id for g in result.scalars().all()} # 3. Determine adds + removals groups_to_add = target_group_ids - existing_group_ids @@ -515,13 +531,15 @@ class GroupTable: # 4. Remove in one bulk delete if groups_to_remove: - db.query(GroupMember).filter( - GroupMember.user_id == user_id, - GroupMember.group_id.in_(groups_to_remove), - ).delete(synchronize_session=False) + await db.execute( + delete(GroupMember).filter( + GroupMember.user_id == user_id, + GroupMember.group_id.in_(groups_to_remove), + ) + ) - db.query(Group).filter(Group.id.in_(groups_to_remove)).update( - {'updated_at': now}, synchronize_session=False + await db.execute( + update(Group).filter(Group.id.in_(groups_to_remove)).values(updated_at=now) ) # 5. Bulk insert missing memberships @@ -537,27 +555,28 @@ class GroupTable: ) if groups_to_add: - db.query(Group).filter(Group.id.in_(groups_to_add)).update( - {'updated_at': now}, synchronize_session=False + await db.execute( + update(Group).filter(Group.id.in_(groups_to_add)).values(updated_at=now) ) - db.commit() + await db.commit() return True except Exception as e: log.exception(e) - db.rollback() + await db.rollback() return False - def add_users_to_group( + async def add_users_to_group( self, id: str, user_ids: Optional[list[str]] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[GroupModel]: try: - with get_db_context(db) as db: - group = db.query(Group).filter_by(id=id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(Group).filter_by(id=id)) + group = result.scalars().first() if not group: return None @@ -574,15 +593,14 @@ class GroupTable: updated_at=now, ) ) - db.flush() # Detect unique constraint violation early + await db.flush() # Detect unique constraint violation early except Exception: - db.rollback() # Clear failed INSERT - db.begin() # Start a new transaction + await db.rollback() # Clear failed INSERT continue # Duplicate → ignore group.updated_at = now - db.commit() - db.refresh(group) + await db.commit() + await db.refresh(group) return GroupModel.model_validate(group) @@ -590,15 +608,16 @@ class GroupTable: log.exception(e) return None - def remove_users_from_group( + async def remove_users_from_group( self, id: str, user_ids: Optional[list[str]] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[GroupModel]: try: - with get_db_context(db) as db: - group = db.query(Group).filter_by(id=id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(Group).filter_by(id=id)) + group = result.scalars().first() if not group: return None @@ -606,15 +625,15 @@ class GroupTable: return GroupModel.model_validate(group) # Remove users from group_member in batch - db.query(GroupMember).filter(GroupMember.group_id == id, GroupMember.user_id.in_(user_ids)).delete( - synchronize_session=False + await db.execute( + delete(GroupMember).filter(GroupMember.group_id == id, GroupMember.user_id.in_(user_ids)) ) # Update group timestamp group.updated_at = int(time.time()) - db.commit() - db.refresh(group) + await db.commit() + await db.refresh(group) return GroupModel.model_validate(group) except Exception as e: diff --git a/backend/open_webui/models/knowledge.py b/backend/open_webui/models/knowledge.py index 30510221fb..68cee36c20 100644 --- a/backend/open_webui/models/knowledge.py +++ b/backend/open_webui/models/knowledge.py @@ -4,8 +4,9 @@ import time from typing import Optional import uuid -from sqlalchemy.orm import Session -from open_webui.internal.db import Base, JSONField, get_db, get_db_context +from sqlalchemy import select, delete, update, or_, func +from sqlalchemy.ext.asyncio import AsyncSession +from open_webui.internal.db import Base, JSONField, get_async_db_context from open_webui.models.files import ( File, @@ -27,7 +28,6 @@ from sqlalchemy import ( Text, JSON, UniqueConstraint, - or_, ) log = logging.getLogger(__name__) @@ -134,25 +134,25 @@ class KnowledgeFileListResponse(BaseModel): class KnowledgeTable: - def _get_access_grants(self, knowledge_id: str, db: Optional[Session] = None) -> list[AccessGrantModel]: - return AccessGrants.get_grants_by_resource('knowledge', knowledge_id, db=db) + async def _get_access_grants(self, knowledge_id: str, db: Optional[AsyncSession] = None) -> list[AccessGrantModel]: + return await AccessGrants.get_grants_by_resource('knowledge', knowledge_id, db=db) - def _to_knowledge_model( + async def _to_knowledge_model( self, knowledge: Knowledge, access_grants: Optional[list[AccessGrantModel]] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> KnowledgeModel: knowledge_data = KnowledgeModel.model_validate(knowledge).model_dump(exclude={'access_grants'}) knowledge_data['access_grants'] = ( - access_grants if access_grants is not None else self._get_access_grants(knowledge_data['id'], db=db) + access_grants if access_grants is not None else await self._get_access_grants(knowledge_data['id'], db=db) ) return KnowledgeModel.model_validate(knowledge_data) - def insert_new_knowledge( - self, user_id: str, form_data: KnowledgeForm, db: Optional[Session] = None + async def insert_new_knowledge( + self, user_id: str, form_data: KnowledgeForm, db: Optional[AsyncSession] = None ) -> Optional[KnowledgeModel]: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: knowledge = KnowledgeModel( **{ **form_data.model_dump(exclude={'access_grants'}), @@ -167,27 +167,28 @@ class KnowledgeTable: try: result = Knowledge(**knowledge.model_dump(exclude={'access_grants'})) db.add(result) - db.commit() - db.refresh(result) - AccessGrants.set_access_grants('knowledge', result.id, form_data.access_grants, db=db) + await db.commit() + await db.refresh(result) + await AccessGrants.set_access_grants('knowledge', result.id, form_data.access_grants, db=db) if result: - return self._to_knowledge_model(result, db=db) + return await self._to_knowledge_model(result, db=db) else: return None except Exception: return None - def get_knowledge_bases( - self, skip: int = 0, limit: int = 30, db: Optional[Session] = None + async def get_knowledge_bases( + self, skip: int = 0, limit: int = 30, db: Optional[AsyncSession] = None ) -> list[KnowledgeUserModel]: - with get_db_context(db) as db: - all_knowledge = db.query(Knowledge).order_by(Knowledge.updated_at.desc()).all() + async with get_async_db_context(db) as db: + result = await db.execute(select(Knowledge).order_by(Knowledge.updated_at.desc())) + all_knowledge = result.scalars().all() user_ids = list(set(knowledge.user_id for knowledge in all_knowledge)) knowledge_ids = [knowledge.id for knowledge in all_knowledge] - users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else [] + users = await Users.get_users_by_user_ids(user_ids, db=db) if user_ids else [] users_dict = {user.id: user for user in users} - grants_map = AccessGrants.get_grants_by_resources('knowledge', knowledge_ids, db=db) + grants_map = await AccessGrants.get_grants_by_resources('knowledge', knowledge_ids, db=db) knowledge_bases = [] for knowledge in all_knowledge: @@ -195,33 +196,33 @@ class KnowledgeTable: knowledge_bases.append( KnowledgeUserModel.model_validate( { - **self._to_knowledge_model( + **(await self._to_knowledge_model( knowledge, access_grants=grants_map.get(knowledge.id, []), db=db, - ).model_dump(), + )).model_dump(), 'user': user.model_dump() if user else None, } ) ) return knowledge_bases - def search_knowledge_bases( + async def search_knowledge_bases( self, user_id: str, filter: dict, skip: int = 0, limit: int = 30, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> KnowledgeListResponse: try: - with get_db_context(db) as db: - query = db.query(Knowledge, User).outerjoin(User, User.id == Knowledge.user_id) + async with get_async_db_context(db) as db: + stmt = select(Knowledge, User).outerjoin(User, User.id == Knowledge.user_id) if filter: query_key = filter.get('query') if query_key: - query = query.filter( + stmt = stmt.filter( or_( Knowledge.name.ilike(f'%{query_key}%'), Knowledge.description.ilike(f'%{query_key}%'), @@ -233,42 +234,46 @@ class KnowledgeTable: view_option = filter.get('view_option') if view_option == 'created': - query = query.filter(Knowledge.user_id == user_id) + stmt = stmt.filter(Knowledge.user_id == user_id) elif view_option == 'shared': - query = query.filter(Knowledge.user_id != user_id) + stmt = stmt.filter(Knowledge.user_id != user_id) - query = AccessGrants.has_permission_filter( + stmt = AccessGrants.has_permission_filter( db=db, - query=query, + query=stmt, DocumentModel=Knowledge, filter=filter, resource_type='knowledge', permission='read', ) - query = query.order_by(Knowledge.updated_at.desc(), Knowledge.id.asc()) + stmt = stmt.order_by(Knowledge.updated_at.desc(), Knowledge.id.asc()) - total = query.count() + count_result = await db.execute( + select(func.count()).select_from(stmt.subquery()) + ) + total = count_result.scalar() if skip: - query = query.offset(skip) + stmt = stmt.offset(skip) if limit: - query = query.limit(limit) + stmt = stmt.limit(limit) - items = query.all() + result = await db.execute(stmt) + items = result.all() knowledge_ids = [kb.id for kb, _ in items] - grants_map = AccessGrants.get_grants_by_resources('knowledge', knowledge_ids, db=db) + grants_map = await AccessGrants.get_grants_by_resources('knowledge', knowledge_ids, db=db) knowledge_bases = [] for knowledge_base, user in items: knowledge_bases.append( KnowledgeUserModel.model_validate( { - **self._to_knowledge_model( + **(await self._to_knowledge_model( knowledge_base, access_grants=grants_map.get(knowledge_base.id, []), db=db, - ).model_dump(), + )).model_dump(), 'user': (UserModel.model_validate(user).model_dump() if user else None), } ) @@ -279,28 +284,27 @@ class KnowledgeTable: print(e) return KnowledgeListResponse(items=[], total=0) - def search_knowledge_files( - self, filter: dict, skip: int = 0, limit: int = 30, db: Optional[Session] = None + async def search_knowledge_files( + self, filter: dict, skip: int = 0, limit: int = 30, db: Optional[AsyncSession] = None ) -> KnowledgeFileListResponse: """ Scalable version: search files across all knowledge bases the user has READ access to, without loading all KBs or using large IN() lists. """ try: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: # Base query: join Knowledge → KnowledgeFile → File - query = ( - db.query(File, User, Knowledge) + stmt = ( + select(File, User, Knowledge) .join(KnowledgeFile, File.id == KnowledgeFile.file_id) .join(Knowledge, KnowledgeFile.knowledge_id == Knowledge.id) .outerjoin(User, User.id == KnowledgeFile.user_id) ) # Apply access-control directly to the joined query - # This makes the database handle filtering, even with 10k+ KBs - query = AccessGrants.has_permission_filter( + stmt = AccessGrants.has_permission_filter( db=db, - query=query, + query=stmt, DocumentModel=Knowledge, filter=filter, resource_type='knowledge', @@ -311,20 +315,24 @@ class KnowledgeTable: if filter: q = filter.get('query') if q: - query = query.filter(File.filename.ilike(f'%{q}%')) + stmt = stmt.filter(File.filename.ilike(f'%{q}%')) # Order by file changes - query = query.order_by(File.updated_at.desc(), File.id.asc()) + stmt = stmt.order_by(File.updated_at.desc(), File.id.asc()) # Count before pagination - total = query.count() + count_result = await db.execute( + select(func.count()).select_from(stmt.subquery()) + ) + total = count_result.scalar() if skip: - query = query.offset(skip) + stmt = stmt.offset(skip) if limit: - query = query.limit(limit) + stmt = stmt.limit(limit) - rows = query.all() + result = await db.execute(stmt) + rows = result.all() items = [] for file, user, knowledge in rows: @@ -332,7 +340,7 @@ class KnowledgeTable: FileUserResponse( **FileModel.model_validate(file).model_dump(), user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None), - collection=self._to_knowledge_model(knowledge, db=db).model_dump(), + collection=(await self._to_knowledge_model(knowledge, db=db)).model_dump(), ) ) @@ -342,14 +350,15 @@ class KnowledgeTable: print('search_knowledge_files error:', e) return KnowledgeFileListResponse(items=[], total=0) - def check_access_by_user_id(self, id, user_id, permission='write', db: Optional[Session] = None) -> bool: - knowledge = self.get_knowledge_by_id(id, db=db) + async def check_access_by_user_id(self, id, user_id, permission='write', db: Optional[AsyncSession] = None) -> bool: + knowledge = await self.get_knowledge_by_id(id, db=db) if not knowledge: return False if knowledge.user_id == user_id: return True - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id, db=db)} - return AccessGrants.has_access( + user_groups = await Groups.get_groups_by_member_id(user_id, db=db) + user_group_ids = {group.id for group in user_groups} + return await AccessGrants.has_access( user_id=user_id, resource_type='knowledge', resource_id=knowledge.id, @@ -358,45 +367,50 @@ class KnowledgeTable: db=db, ) - def get_knowledge_bases_by_user_id( - self, user_id: str, permission: str = 'write', db: Optional[Session] = None + async def get_knowledge_bases_by_user_id( + self, user_id: str, permission: str = 'write', db: Optional[AsyncSession] = None ) -> list[KnowledgeUserModel]: - knowledge_bases = self.get_knowledge_bases(db=db) - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id, db=db)} - return [ - knowledge_base - for knowledge_base in knowledge_bases - if knowledge_base.user_id == user_id - or AccessGrants.has_access( + knowledge_bases = await self.get_knowledge_bases(db=db) + user_groups = await Groups.get_groups_by_member_id(user_id, db=db) + user_group_ids = {group.id for group in user_groups} + + result = [] + for knowledge_base in knowledge_bases: + if knowledge_base.user_id == user_id: + result.append(knowledge_base) + elif await AccessGrants.has_access( user_id=user_id, resource_type='knowledge', resource_id=knowledge_base.id, permission=permission, user_group_ids=user_group_ids, db=db, - ) - ] + ): + result.append(knowledge_base) + return result - def get_knowledge_by_id(self, id: str, db: Optional[Session] = None) -> Optional[KnowledgeModel]: + async def get_knowledge_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[KnowledgeModel]: try: - with get_db_context(db) as db: - knowledge = db.query(Knowledge).filter_by(id=id).first() - return self._to_knowledge_model(knowledge, db=db) if knowledge else None + async with get_async_db_context(db) as db: + result = await db.execute(select(Knowledge).filter_by(id=id)) + knowledge = result.scalars().first() + return await self._to_knowledge_model(knowledge, db=db) if knowledge else None except Exception: return None - def get_knowledge_by_id_and_user_id( - self, id: str, user_id: str, db: Optional[Session] = None + async def get_knowledge_by_id_and_user_id( + self, id: str, user_id: str, db: Optional[AsyncSession] = None ) -> Optional[KnowledgeModel]: - knowledge = self.get_knowledge_by_id(id, db=db) + knowledge = await self.get_knowledge_by_id(id, db=db) if not knowledge: return None if knowledge.user_id == user_id: return knowledge - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id, db=db)} - if AccessGrants.has_access( + user_groups = await Groups.get_groups_by_member_id(user_id, db=db) + user_group_ids = {group.id for group in user_groups} + if await AccessGrants.has_access( user_id=user_id, resource_type='knowledge', resource_id=knowledge.id, @@ -407,19 +421,19 @@ class KnowledgeTable: return knowledge return None - def get_knowledges_by_file_id(self, file_id: str, db: Optional[Session] = None) -> list[KnowledgeModel]: + async def get_knowledges_by_file_id(self, file_id: str, db: Optional[AsyncSession] = None) -> list[KnowledgeModel]: try: - with get_db_context(db) as db: - knowledges = ( - db.query(Knowledge) + async with get_async_db_context(db) as db: + result = await db.execute( + select(Knowledge) .join(KnowledgeFile, Knowledge.id == KnowledgeFile.knowledge_id) .filter(KnowledgeFile.file_id == file_id) - .all() ) + knowledges = result.scalars().all() knowledge_ids = [k.id for k in knowledges] - grants_map = AccessGrants.get_grants_by_resources('knowledge', knowledge_ids, db=db) + grants_map = await AccessGrants.get_grants_by_resources('knowledge', knowledge_ids, db=db) return [ - self._to_knowledge_model( + await self._to_knowledge_model( knowledge, access_grants=grants_map.get(knowledge.id, []), db=db, @@ -429,19 +443,19 @@ class KnowledgeTable: except Exception: return [] - def search_files_by_id( + async def search_files_by_id( self, knowledge_id: str, user_id: str, filter: dict, skip: int = 0, limit: int = 30, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> KnowledgeFileListResponse: try: - with get_db_context(db) as db: - query = ( - db.query(File, User) + async with get_async_db_context(db) as db: + stmt = ( + select(File, User) .join(KnowledgeFile, File.id == KnowledgeFile.file_id) .outerjoin(User, User.id == KnowledgeFile.user_id) .filter(KnowledgeFile.knowledge_id == knowledge_id) @@ -453,13 +467,13 @@ class KnowledgeTable: if filter: query_key = filter.get('query') if query_key: - query = query.filter(or_(File.filename.ilike(f'%{query_key}%'))) + stmt = stmt.filter(or_(File.filename.ilike(f'%{query_key}%'))) view_option = filter.get('view_option') if view_option == 'created': - query = query.filter(KnowledgeFile.user_id == user_id) + stmt = stmt.filter(KnowledgeFile.user_id == user_id) elif view_option == 'shared': - query = query.filter(KnowledgeFile.user_id != user_id) + stmt = stmt.filter(KnowledgeFile.user_id != user_id) order_by = filter.get('order_by') direction = filter.get('direction') @@ -473,17 +487,21 @@ class KnowledgeTable: primary_sort = File.updated_at.asc() if is_asc else File.updated_at.desc() # Apply sort with secondary key for deterministic pagination - query = query.order_by(primary_sort, File.id.asc()) + stmt = stmt.order_by(primary_sort, File.id.asc()) # Count BEFORE pagination - total = query.count() + count_result = await db.execute( + select(func.count()).select_from(stmt.subquery()) + ) + total = count_result.scalar() if skip: - query = query.offset(skip) + stmt = stmt.offset(skip) if limit: - query = query.limit(limit) + stmt = stmt.limit(limit) - items = query.all() + result = await db.execute(stmt) + items = result.all() files = [] for file, user in items: @@ -499,35 +517,34 @@ class KnowledgeTable: print(e) return KnowledgeFileListResponse(items=[], total=0) - def get_files_by_id(self, knowledge_id: str, db: Optional[Session] = None) -> list[FileModel]: + async def get_files_by_id(self, knowledge_id: str, db: Optional[AsyncSession] = None) -> list[FileModel]: try: - with get_db_context(db) as db: - files = ( - db.query(File) + async with get_async_db_context(db) as db: + result = await db.execute( + select(File) .join(KnowledgeFile, File.id == KnowledgeFile.file_id) .filter(KnowledgeFile.knowledge_id == knowledge_id) - .all() ) + files = result.scalars().all() return [FileModel.model_validate(file) for file in files] except Exception: return [] - def get_file_metadatas_by_id(self, knowledge_id: str, db: Optional[Session] = None) -> list[FileMetadataResponse]: + async def get_file_metadatas_by_id(self, knowledge_id: str, db: Optional[AsyncSession] = None) -> list[FileMetadataResponse]: try: - with get_db_context(db) as db: - files = self.get_files_by_id(knowledge_id, db=db) - return [FileMetadataResponse(**file.model_dump()) for file in files] + files = await self.get_files_by_id(knowledge_id, db=db) + return [FileMetadataResponse(**file.model_dump()) for file in files] except Exception: return [] - def add_file_to_knowledge_by_id( + async def add_file_to_knowledge_by_id( self, knowledge_id: str, file_id: str, user_id: str, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[KnowledgeFileModel]: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: knowledge_file = KnowledgeFileModel( **{ 'id': str(uuid.uuid4()), @@ -542,8 +559,8 @@ class KnowledgeTable: try: result = KnowledgeFile(**knowledge_file.model_dump()) db.add(result) - db.commit() - db.refresh(result) + await db.commit() + await db.refresh(result) if result: return KnowledgeFileModel.model_validate(result) else: @@ -551,103 +568,103 @@ class KnowledgeTable: except Exception: return None - def has_file(self, knowledge_id: str, file_id: str, db: Optional[Session] = None) -> bool: + async def has_file(self, knowledge_id: str, file_id: str, db: Optional[AsyncSession] = None) -> bool: """Check whether a file belongs to a knowledge base.""" try: - with get_db_context(db) as db: - return db.query(KnowledgeFile).filter_by(knowledge_id=knowledge_id, file_id=file_id).first() is not None + async with get_async_db_context(db) as db: + result = await db.execute( + select(KnowledgeFile).filter_by(knowledge_id=knowledge_id, file_id=file_id).limit(1) + ) + return result.scalars().first() is not None except Exception: return False - def remove_file_from_knowledge_by_id(self, knowledge_id: str, file_id: str, db: Optional[Session] = None) -> bool: + async def remove_file_from_knowledge_by_id(self, knowledge_id: str, file_id: str, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: - db.query(KnowledgeFile).filter_by(knowledge_id=knowledge_id, file_id=file_id).delete() - db.commit() + async with get_async_db_context(db) as db: + await db.execute(delete(KnowledgeFile).filter_by(knowledge_id=knowledge_id, file_id=file_id)) + await db.commit() return True except Exception: return False - def reset_knowledge_by_id(self, id: str, db: Optional[Session] = None) -> Optional[KnowledgeModel]: + async def reset_knowledge_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[KnowledgeModel]: try: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: # Delete all knowledge_file entries for this knowledge_id - db.query(KnowledgeFile).filter_by(knowledge_id=id).delete() - db.commit() + await db.execute(delete(KnowledgeFile).filter_by(knowledge_id=id)) + await db.commit() # Update the knowledge entry's updated_at timestamp - db.query(Knowledge).filter_by(id=id).update( - { - 'updated_at': int(time.time()), - } + await db.execute( + update(Knowledge).filter_by(id=id).values(updated_at=int(time.time())) ) - db.commit() + await db.commit() - return self.get_knowledge_by_id(id=id, db=db) + return await self.get_knowledge_by_id(id=id, db=db) except Exception as e: log.exception(e) return None - def update_knowledge_by_id( + async def update_knowledge_by_id( self, id: str, form_data: KnowledgeForm, overwrite: bool = False, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[KnowledgeModel]: try: - with get_db_context(db) as db: - knowledge = self.get_knowledge_by_id(id=id, db=db) - db.query(Knowledge).filter_by(id=id).update( - { + async with get_async_db_context(db) as db: + await db.execute( + update(Knowledge).filter_by(id=id).values( **form_data.model_dump(exclude={'access_grants'}), - 'updated_at': int(time.time()), - } + updated_at=int(time.time()), + ) ) - db.commit() + await db.commit() if form_data.access_grants is not None: - AccessGrants.set_access_grants('knowledge', id, form_data.access_grants, db=db) - return self.get_knowledge_by_id(id=id, db=db) + await AccessGrants.set_access_grants('knowledge', id, form_data.access_grants, db=db) + return await self.get_knowledge_by_id(id=id, db=db) except Exception as e: log.exception(e) return None - def update_knowledge_data_by_id( - self, id: str, data: dict, db: Optional[Session] = None + async def update_knowledge_data_by_id( + self, id: str, data: dict, db: Optional[AsyncSession] = None ) -> Optional[KnowledgeModel]: try: - with get_db_context(db) as db: - knowledge = self.get_knowledge_by_id(id=id, db=db) - db.query(Knowledge).filter_by(id=id).update( - { - 'data': data, - 'updated_at': int(time.time()), - } + async with get_async_db_context(db) as db: + await db.execute( + update(Knowledge).filter_by(id=id).values( + data=data, + updated_at=int(time.time()), + ) ) - db.commit() - return self.get_knowledge_by_id(id=id, db=db) + await db.commit() + return await self.get_knowledge_by_id(id=id, db=db) except Exception as e: log.exception(e) return None - def delete_knowledge_by_id(self, id: str, db: Optional[Session] = None) -> bool: + async def delete_knowledge_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: - AccessGrants.revoke_all_access('knowledge', id, db=db) - db.query(Knowledge).filter_by(id=id).delete() - db.commit() + async with get_async_db_context(db) as db: + await AccessGrants.revoke_all_access('knowledge', id, db=db) + await db.execute(delete(Knowledge).filter_by(id=id)) + await db.commit() return True except Exception: return False - def delete_all_knowledge(self, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: + async def delete_all_knowledge(self, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: try: - knowledge_ids = [row[0] for row in db.query(Knowledge.id).all()] + result = await db.execute(select(Knowledge.id)) + knowledge_ids = [row[0] for row in result.all()] for knowledge_id in knowledge_ids: - AccessGrants.revoke_all_access('knowledge', knowledge_id, db=db) - db.query(Knowledge).delete() - db.commit() + await AccessGrants.revoke_all_access('knowledge', knowledge_id, db=db) + await db.execute(delete(Knowledge)) + await db.commit() return True except Exception: diff --git a/backend/open_webui/models/memories.py b/backend/open_webui/models/memories.py index 7c34de9f07..e956826800 100644 --- a/backend/open_webui/models/memories.py +++ b/backend/open_webui/models/memories.py @@ -2,8 +2,9 @@ import time import uuid from typing import Optional -from sqlalchemy.orm import Session -from open_webui.internal.db import Base, get_db, get_db_context +from sqlalchemy import select, delete +from sqlalchemy.ext.asyncio import AsyncSession +from open_webui.internal.db import Base, get_async_db_context from pydantic import BaseModel, ConfigDict from sqlalchemy import BigInteger, Column, String, Text @@ -40,13 +41,13 @@ class MemoryModel(BaseModel): class MemoriesTable: - def insert_new_memory( + async def insert_new_memory( self, user_id: str, content: str, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[MemoryModel]: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: id = str(uuid.uuid4()) memory = MemoryModel( @@ -60,90 +61,92 @@ class MemoriesTable: ) result = Memory(**memory.model_dump()) db.add(result) - db.commit() - db.refresh(result) + await db.commit() + await db.refresh(result) if result: return MemoryModel.model_validate(result) else: return None - def update_memory_by_id_and_user_id( + async def update_memory_by_id_and_user_id( self, id: str, user_id: str, content: str, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[MemoryModel]: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: try: - memory = db.get(Memory, id) + memory = await db.get(Memory, id) if not memory or memory.user_id != user_id: return None memory.content = content memory.updated_at = int(time.time()) - db.commit() - db.refresh(memory) + await db.commit() + await db.refresh(memory) return MemoryModel.model_validate(memory) except Exception: return None - def get_memories(self, db: Optional[Session] = None) -> list[MemoryModel]: - with get_db_context(db) as db: + async def get_memories(self, db: Optional[AsyncSession] = None) -> list[MemoryModel]: + async with get_async_db_context(db) as db: try: - memories = db.query(Memory).all() + result = await db.execute(select(Memory)) + memories = result.scalars().all() return [MemoryModel.model_validate(memory) for memory in memories] except Exception: return None - def get_memories_by_user_id(self, user_id: str, db: Optional[Session] = None) -> list[MemoryModel]: - with get_db_context(db) as db: + async def get_memories_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[MemoryModel]: + async with get_async_db_context(db) as db: try: - memories = db.query(Memory).filter_by(user_id=user_id).all() + result = await db.execute(select(Memory).filter_by(user_id=user_id)) + memories = result.scalars().all() return [MemoryModel.model_validate(memory) for memory in memories] except Exception: return None - def get_memory_by_id(self, id: str, db: Optional[Session] = None) -> Optional[MemoryModel]: - with get_db_context(db) as db: + async def get_memory_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[MemoryModel]: + async with get_async_db_context(db) as db: try: - memory = db.get(Memory, id) - return MemoryModel.model_validate(memory) + memory = await db.get(Memory, id) + return MemoryModel.model_validate(memory) if memory else None except Exception: return None - def delete_memory_by_id(self, id: str, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: + async def delete_memory_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: try: - db.query(Memory).filter_by(id=id).delete() - db.commit() + await db.execute(delete(Memory).filter_by(id=id)) + await db.commit() return True except Exception: return False - def delete_memories_by_user_id(self, user_id: str, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: + async def delete_memories_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: try: - db.query(Memory).filter_by(user_id=user_id).delete() - db.commit() + await db.execute(delete(Memory).filter_by(user_id=user_id)) + await db.commit() return True except Exception: return False - def delete_memory_by_id_and_user_id(self, id: str, user_id: str, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: + async def delete_memory_by_id_and_user_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: try: - memory = db.get(Memory, id) + memory = await db.get(Memory, id) if not memory or memory.user_id != user_id: return None # Delete the memory - db.delete(memory) - db.commit() + await db.delete(memory) + await db.commit() return True except Exception: diff --git a/backend/open_webui/models/messages.py b/backend/open_webui/models/messages.py index 034eaac160..c9af45ebf5 100644 --- a/backend/open_webui/models/messages.py +++ b/backend/open_webui/models/messages.py @@ -3,8 +3,9 @@ import time import uuid from typing import Optional -from sqlalchemy.orm import Session -from open_webui.internal.db import Base, JSONField, get_db, get_db_context +from sqlalchemy import select, delete, func +from sqlalchemy.ext.asyncio import AsyncSession +from open_webui.internal.db import Base, JSONField, get_async_db_context from open_webui.models.tags import TagModel, Tag, Tags from open_webui.models.users import Users, User, UserNameResponse from open_webui.models.channels import Channels, ChannelMember @@ -12,7 +13,7 @@ from open_webui.models.channels import Channels, ChannelMember from pydantic import BaseModel, ConfigDict, field_validator from sqlalchemy import BigInteger, Boolean, Column, String, Text, JSON -from sqlalchemy import or_, func, select, and_, text +from sqlalchemy import or_, func, and_, text from sqlalchemy.sql import exists #################### @@ -137,15 +138,15 @@ class MessageResponse(MessageReplyToResponse): class MessageTable: - def insert_new_message( + async def insert_new_message( self, form_data: MessageForm, channel_id: str, user_id: str, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[MessageModel]: - with get_db_context(db) as db: - channel_member = Channels.join_channel(channel_id, user_id) + async with get_async_db_context(db) as db: + channel_member = await Channels.join_channel(channel_id, user_id) id = str(uuid.uuid4()) ts = int(time.time_ns()) @@ -170,38 +171,38 @@ class MessageTable: result = Message(**message.model_dump()) db.add(result) - db.commit() - db.refresh(result) + await db.commit() + await db.refresh(result) return MessageModel.model_validate(result) if result else None - def get_message_by_id( + async def get_message_by_id( self, id: str, include_thread_replies: Optional[bool] = True, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[MessageResponse]: - with get_db_context(db) as db: - message = db.get(Message, id) + async with get_async_db_context(db) as db: + message = await db.get(Message, id) if not message: return None reply_to_message = ( - self.get_message_by_id(message.reply_to_id, include_thread_replies=False, db=db) + await self.get_message_by_id(message.reply_to_id, include_thread_replies=False, db=db) if message.reply_to_id else None ) - reactions = self.get_reactions_by_message_id(id, db=db) + reactions = await self.get_reactions_by_message_id(id, db=db) thread_replies = [] if include_thread_replies: - thread_replies = self.get_thread_replies_by_message_id(id, db=db) + thread_replies = await self.get_thread_replies_by_message_id(id, db=db) # Check if message was sent by webhook (webhook info in meta takes precedence) webhook_info = message.meta.get('webhook') if message.meta else None if webhook_info and webhook_info.get('id'): # Look up webhook by ID to get current name - webhook = Channels.get_webhook_by_id(webhook_info.get('id'), db=db) + webhook = await Channels.get_webhook_by_id(webhook_info.get('id'), db=db) if webhook: user_info = { 'id': webhook.id, @@ -216,7 +217,7 @@ class MessageTable: 'role': 'webhook', } else: - user = Users.get_user_by_id(message.user_id, db=db) + user = await Users.get_user_by_id(message.user_id, db=db) user_info = user.model_dump() if user else None return MessageResponse.model_validate( @@ -230,34 +231,41 @@ class MessageTable: } ) - def get_thread_replies_by_message_id(self, id: str, db: Optional[Session] = None) -> list[MessageReplyToResponse]: - with get_db_context(db) as db: - all_messages = db.query(Message).filter_by(parent_id=id).order_by(Message.created_at.desc()).all() + async def _resolve_user_info(self, message: Message, db: AsyncSession) -> Optional[dict]: + """Resolve user info from message, handling webhook messages.""" + webhook_info = message.meta.get('webhook') if message.meta else None + if webhook_info and webhook_info.get('id'): + webhook = await Channels.get_webhook_by_id(webhook_info.get('id'), db=db) + if webhook: + return { + 'id': webhook.id, + 'name': webhook.name, + 'role': 'webhook', + } + else: + return { + 'id': webhook_info.get('id'), + 'name': 'Deleted Webhook', + 'role': 'webhook', + } + return None + + async def get_thread_replies_by_message_id(self, id: str, db: Optional[AsyncSession] = None) -> list[MessageReplyToResponse]: + async with get_async_db_context(db) as db: + result = await db.execute( + select(Message).filter_by(parent_id=id).order_by(Message.created_at.desc()) + ) + all_messages = result.scalars().all() messages = [] for message in all_messages: reply_to_message = ( - self.get_message_by_id(message.reply_to_id, include_thread_replies=False, db=db) + await self.get_message_by_id(message.reply_to_id, include_thread_replies=False, db=db) if message.reply_to_id else None ) - webhook_info = message.meta.get('webhook') if message.meta else None - user_info = None - if webhook_info and webhook_info.get('id'): - webhook = Channels.get_webhook_by_id(webhook_info.get('id'), db=db) - if webhook: - user_info = { - 'id': webhook.id, - 'name': webhook.name, - 'role': 'webhook', - } - else: - user_info = { - 'id': webhook_info.get('id'), - 'name': 'Deleted Webhook', - 'role': 'webhook', - } + user_info = await self._resolve_user_info(message, db) messages.append( MessageReplyToResponse.model_validate( @@ -270,51 +278,37 @@ class MessageTable: ) return messages - def get_reply_user_ids_by_message_id(self, id: str, db: Optional[Session] = None) -> list[str]: - with get_db_context(db) as db: - return [message.user_id for message in db.query(Message).filter_by(parent_id=id).all()] + async def get_reply_user_ids_by_message_id(self, id: str, db: Optional[AsyncSession] = None) -> list[str]: + async with get_async_db_context(db) as db: + result = await db.execute(select(Message.user_id).filter_by(parent_id=id)) + return [row[0] for row in result.all()] - def get_messages_by_channel_id( + async def get_messages_by_channel_id( self, channel_id: str, skip: int = 0, limit: int = 50, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[MessageReplyToResponse]: - with get_db_context(db) as db: - all_messages = ( - db.query(Message) + async with get_async_db_context(db) as db: + result = await db.execute( + select(Message) .filter_by(channel_id=channel_id, parent_id=None) .order_by(Message.created_at.desc()) .offset(skip) .limit(limit) - .all() ) + all_messages = result.scalars().all() messages = [] for message in all_messages: reply_to_message = ( - self.get_message_by_id(message.reply_to_id, include_thread_replies=False, db=db) + await self.get_message_by_id(message.reply_to_id, include_thread_replies=False, db=db) if message.reply_to_id else None ) - webhook_info = message.meta.get('webhook') if message.meta else None - user_info = None - if webhook_info and webhook_info.get('id'): - webhook = Channels.get_webhook_by_id(webhook_info.get('id'), db=db) - if webhook: - user_info = { - 'id': webhook.id, - 'name': webhook.name, - 'role': 'webhook', - } - else: - user_info = { - 'id': webhook_info.get('id'), - 'name': 'Deleted Webhook', - 'role': 'webhook', - } + user_info = await self._resolve_user_info(message, db) messages.append( MessageReplyToResponse.model_validate( @@ -327,28 +321,28 @@ class MessageTable: ) return messages - def get_messages_by_parent_id( + async def get_messages_by_parent_id( self, channel_id: str, parent_id: str, skip: int = 0, limit: int = 50, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[MessageReplyToResponse]: - with get_db_context(db) as db: - message = db.get(Message, parent_id) + async with get_async_db_context(db) as db: + message = await db.get(Message, parent_id) if not message: return [] - all_messages = ( - db.query(Message) + result = await db.execute( + select(Message) .filter_by(channel_id=channel_id, parent_id=parent_id) .order_by(Message.created_at.desc()) .offset(skip) .limit(limit) - .all() ) + all_messages = list(result.scalars().all()) # If length of all_messages is less than limit, then add the parent message if len(all_messages) < limit: @@ -357,27 +351,12 @@ class MessageTable: messages = [] for message in all_messages: reply_to_message = ( - self.get_message_by_id(message.reply_to_id, include_thread_replies=False, db=db) + await self.get_message_by_id(message.reply_to_id, include_thread_replies=False, db=db) if message.reply_to_id else None ) - webhook_info = message.meta.get('webhook') if message.meta else None - user_info = None - if webhook_info and webhook_info.get('id'): - webhook = Channels.get_webhook_by_id(webhook_info.get('id'), db=db) - if webhook: - user_info = { - 'id': webhook.id, - 'name': webhook.name, - 'role': 'webhook', - } - else: - user_info = { - 'id': webhook_info.get('id'), - 'name': 'Deleted Webhook', - 'role': 'webhook', - } + user_info = await self._resolve_user_info(message, db) messages.append( MessageReplyToResponse.model_validate( @@ -390,34 +369,37 @@ class MessageTable: ) return messages - def get_last_message_by_channel_id(self, channel_id: str, db: Optional[Session] = None) -> Optional[MessageModel]: - with get_db_context(db) as db: - message = db.query(Message).filter_by(channel_id=channel_id).order_by(Message.created_at.desc()).first() + async def get_last_message_by_channel_id(self, channel_id: str, db: Optional[AsyncSession] = None) -> Optional[MessageModel]: + async with get_async_db_context(db) as db: + result = await db.execute( + select(Message).filter_by(channel_id=channel_id).order_by(Message.created_at.desc()).limit(1) + ) + message = result.scalars().first() return MessageModel.model_validate(message) if message else None - def get_pinned_messages_by_channel_id( + async def get_pinned_messages_by_channel_id( self, channel_id: str, skip: int = 0, limit: int = 50, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[MessageModel]: - with get_db_context(db) as db: - all_messages = ( - db.query(Message) + async with get_async_db_context(db) as db: + result = await db.execute( + select(Message) .filter_by(channel_id=channel_id, is_pinned=True) .order_by(Message.pinned_at.desc()) .offset(skip) .limit(limit) - .all() ) + all_messages = result.scalars().all() return [MessageModel.model_validate(message) for message in all_messages] - def update_message_by_id( - self, id: str, form_data: MessageForm, db: Optional[Session] = None + async def update_message_by_id( + self, id: str, form_data: MessageForm, db: Optional[AsyncSession] = None ) -> Optional[MessageModel]: - with get_db_context(db) as db: - message = db.get(Message, id) + async with get_async_db_context(db) as db: + message = await db.get(Message, id) message.content = form_data.content message.data = { **(message.data if message.data else {}), @@ -428,49 +410,53 @@ class MessageTable: **(form_data.meta if form_data.meta else {}), } message.updated_at = int(time.time_ns()) - db.commit() - db.refresh(message) + await db.commit() + await db.refresh(message) return MessageModel.model_validate(message) if message else None - def update_is_pinned_by_id( + async def update_is_pinned_by_id( self, id: str, is_pinned: bool, pinned_by: Optional[str] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[MessageModel]: - with get_db_context(db) as db: - message = db.get(Message, id) + async with get_async_db_context(db) as db: + message = await db.get(Message, id) message.is_pinned = is_pinned message.pinned_at = int(time.time_ns()) if is_pinned else None message.pinned_by = pinned_by if is_pinned else None - db.commit() - db.refresh(message) + await db.commit() + await db.refresh(message) return MessageModel.model_validate(message) if message else None - def get_unread_message_count( + async def get_unread_message_count( self, channel_id: str, user_id: str, last_read_at: Optional[int] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> int: - with get_db_context(db) as db: - query = db.query(Message).filter( + async with get_async_db_context(db) as db: + stmt = select(func.count(Message.id)).filter( Message.channel_id == channel_id, Message.parent_id == None, # only count top-level messages Message.created_at > (last_read_at if last_read_at else 0), ) if user_id: - query = query.filter(Message.user_id != user_id) - return query.count() + stmt = stmt.filter(Message.user_id != user_id) + result = await db.execute(stmt) + return result.scalar() - def add_reaction_to_message( - self, id: str, user_id: str, name: str, db: Optional[Session] = None + async def add_reaction_to_message( + self, id: str, user_id: str, name: str, db: Optional[AsyncSession] = None ) -> Optional[MessageReactionModel]: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: # check for existing reaction - existing_reaction = db.query(MessageReaction).filter_by(message_id=id, user_id=user_id, name=name).first() + result = await db.execute( + select(MessageReaction).filter_by(message_id=id, user_id=user_id, name=name) + ) + existing_reaction = result.scalars().first() if existing_reaction: return MessageReactionModel.model_validate(existing_reaction) @@ -484,19 +470,19 @@ class MessageTable: ) result = MessageReaction(**reaction.model_dump()) db.add(result) - db.commit() - db.refresh(result) + await db.commit() + await db.refresh(result) return MessageReactionModel.model_validate(result) if result else None - def get_reactions_by_message_id(self, id: str, db: Optional[Session] = None) -> list[Reactions]: - with get_db_context(db) as db: + async def get_reactions_by_message_id(self, id: str, db: Optional[AsyncSession] = None) -> list[Reactions]: + async with get_async_db_context(db) as db: # JOIN User so all user info is fetched in one query - results = ( - db.query(MessageReaction, User) + result = await db.execute( + select(MessageReaction, User) .join(User, MessageReaction.user_id == User.id) .filter(MessageReaction.message_id == id) - .all() ) + results = result.all() reactions = {} @@ -518,58 +504,60 @@ class MessageTable: return [Reactions(**reaction) for reaction in reactions.values()] - def remove_reaction_by_id_and_user_id_and_name( - self, id: str, user_id: str, name: str, db: Optional[Session] = None + async def remove_reaction_by_id_and_user_id_and_name( + self, id: str, user_id: str, name: str, db: Optional[AsyncSession] = None ) -> bool: - with get_db_context(db) as db: - db.query(MessageReaction).filter_by(message_id=id, user_id=user_id, name=name).delete() - db.commit() + async with get_async_db_context(db) as db: + await db.execute(delete(MessageReaction).filter_by(message_id=id, user_id=user_id, name=name)) + await db.commit() return True - def delete_reactions_by_id(self, id: str, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: - db.query(MessageReaction).filter_by(message_id=id).delete() - db.commit() + async def delete_reactions_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: + await db.execute(delete(MessageReaction).filter_by(message_id=id)) + await db.commit() return True - def delete_replies_by_id(self, id: str, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: - db.query(Message).filter_by(parent_id=id).delete() - db.commit() + async def delete_replies_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: + await db.execute(delete(Message).filter_by(parent_id=id)) + await db.commit() return True - def delete_message_by_id(self, id: str, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: - db.query(Message).filter_by(id=id).delete() + async def delete_message_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: + await db.execute(delete(Message).filter_by(id=id)) # Delete all reactions to this message - db.query(MessageReaction).filter_by(message_id=id).delete() + await db.execute(delete(MessageReaction).filter_by(message_id=id)) - db.commit() + await db.commit() return True - def search_messages_by_channel_ids( + async def search_messages_by_channel_ids( self, channel_ids: list[str], query: str, start_timestamp: Optional[int] = None, end_timestamp: Optional[int] = None, limit: int = 10, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[MessageModel]: """Search messages in specified channels by content.""" - with get_db_context(db) as db: - query_builder = db.query(Message).filter( + async with get_async_db_context(db) as db: + stmt = select(Message).filter( Message.channel_id.in_(channel_ids), Message.content.ilike(f'%{query}%'), ) if start_timestamp: - query_builder = query_builder.filter(Message.created_at >= start_timestamp) + stmt = stmt.filter(Message.created_at >= start_timestamp) if end_timestamp: - query_builder = query_builder.filter(Message.created_at <= end_timestamp) + stmt = stmt.filter(Message.created_at <= end_timestamp) - messages = query_builder.order_by(Message.created_at.desc()).limit(limit).all() + stmt = stmt.order_by(Message.created_at.desc()).limit(limit) + result = await db.execute(stmt) + messages = result.scalars().all() return [MessageModel.model_validate(msg) for msg in messages] diff --git a/backend/open_webui/models/models.py b/backend/open_webui/models/models.py index 7cab2c830e..4664a71b85 100755 --- a/backend/open_webui/models/models.py +++ b/backend/open_webui/models/models.py @@ -2,8 +2,9 @@ import logging import time from typing import Optional -from sqlalchemy.orm import Session -from open_webui.internal.db import Base, JSONField, get_db, get_db_context +from sqlalchemy import select, delete, update, or_, func, String, cast +from sqlalchemy.ext.asyncio import AsyncSession +from open_webui.internal.db import Base, JSONField, get_async_db_context from open_webui.models.groups import Groups from open_webui.models.users import User, UserModel, Users, UserResponse @@ -12,9 +13,6 @@ from open_webui.models.access_grants import AccessGrantModel, AccessGrants from pydantic import BaseModel, ConfigDict, Field, model_validator -from sqlalchemy import String, cast, or_, and_, func -from sqlalchemy.dialects import postgresql, sqlite - from sqlalchemy.dialects.postgresql import JSONB from sqlalchemy import BigInteger, Column, Text, Boolean @@ -154,26 +152,26 @@ class ModelForm(BaseModel): class ModelsTable: - def _get_access_grants(self, model_id: str, db: Optional[Session] = None) -> list[AccessGrantModel]: - return AccessGrants.get_grants_by_resource('model', model_id, db=db) + async def _get_access_grants(self, model_id: str, db: Optional[AsyncSession] = None) -> list[AccessGrantModel]: + return await AccessGrants.get_grants_by_resource('model', model_id, db=db) - def _to_model_model( + async def _to_model_model( self, model: Model, access_grants: Optional[list[AccessGrantModel]] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> ModelModel: model_data = ModelModel.model_validate(model).model_dump(exclude={'access_grants'}) model_data['access_grants'] = ( - access_grants if access_grants is not None else self._get_access_grants(model_data['id'], db=db) + access_grants if access_grants is not None else await self._get_access_grants(model_data['id'], db=db) ) return ModelModel.model_validate(model_data) - def insert_new_model( - self, form_data: ModelForm, user_id: str, db: Optional[Session] = None + async def insert_new_model( + self, form_data: ModelForm, user_id: str, db: Optional[AsyncSession] = None ) -> Optional[ModelModel]: try: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: result = Model( **{ **form_data.model_dump(exclude={'access_grants'}), @@ -183,37 +181,39 @@ class ModelsTable: } ) db.add(result) - db.commit() - db.refresh(result) - AccessGrants.set_access_grants('model', result.id, form_data.access_grants, db=db) + await db.commit() + await db.refresh(result) + await AccessGrants.set_access_grants('model', result.id, form_data.access_grants, db=db) if result: - return self._to_model_model(result, db=db) + return await self._to_model_model(result, db=db) else: return None except Exception as e: log.exception(f'Failed to insert a new model: {e}') return None - def get_all_models(self, db: Optional[Session] = None) -> list[ModelModel]: - with get_db_context(db) as db: - all_models = db.query(Model).all() + async def get_all_models(self, db: Optional[AsyncSession] = None) -> list[ModelModel]: + async with get_async_db_context(db) as db: + result = await db.execute(select(Model)) + all_models = result.scalars().all() model_ids = [model.id for model in all_models] - grants_map = AccessGrants.get_grants_by_resources('model', model_ids, db=db) + grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db) return [ - self._to_model_model(model, access_grants=grants_map.get(model.id, []), db=db) for model in all_models + await self._to_model_model(model, access_grants=grants_map.get(model.id, []), db=db) for model in all_models ] - def get_models(self, db: Optional[Session] = None) -> list[ModelUserResponse]: - with get_db_context(db) as db: - all_models = db.query(Model).filter(Model.base_model_id != None).all() + async def get_models(self, db: Optional[AsyncSession] = None) -> list[ModelUserResponse]: + async with get_async_db_context(db) as db: + result = await db.execute(select(Model).filter(Model.base_model_id != None)) + all_models = result.scalars().all() user_ids = list(set(model.user_id for model in all_models)) model_ids = [model.id for model in all_models] - users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else [] + users = await Users.get_users_by_user_ids(user_ids, db=db) if user_ids else [] users_dict = {user.id: user for user in users} - grants_map = AccessGrants.get_grants_by_resources('model', model_ids, db=db) + grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db) models = [] for model in all_models: @@ -221,44 +221,48 @@ class ModelsTable: models.append( ModelUserResponse.model_validate( { - **self._to_model_model( + **(await self._to_model_model( model, access_grants=grants_map.get(model.id, []), db=db, - ).model_dump(), + )).model_dump(), 'user': user.model_dump() if user else None, } ) ) return models - def get_base_models(self, db: Optional[Session] = None) -> list[ModelModel]: - with get_db_context(db) as db: - all_models = db.query(Model).filter(Model.base_model_id == None).all() + async def get_base_models(self, db: Optional[AsyncSession] = None) -> list[ModelModel]: + async with get_async_db_context(db) as db: + result = await db.execute(select(Model).filter(Model.base_model_id == None)) + all_models = result.scalars().all() model_ids = [model.id for model in all_models] - grants_map = AccessGrants.get_grants_by_resources('model', model_ids, db=db) + grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db) return [ - self._to_model_model(model, access_grants=grants_map.get(model.id, []), db=db) for model in all_models + await self._to_model_model(model, access_grants=grants_map.get(model.id, []), db=db) for model in all_models ] - def get_models_by_user_id( - self, user_id: str, permission: str = 'write', db: Optional[Session] = None + async def get_models_by_user_id( + self, user_id: str, permission: str = 'write', db: Optional[AsyncSession] = None ) -> list[ModelUserResponse]: - models = self.get_models(db=db) - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id, db=db)} - return [ - model - for model in models - if model.user_id == user_id - or AccessGrants.has_access( + models = await self.get_models(db=db) + user_groups = await Groups.get_groups_by_member_id(user_id, db=db) + user_group_ids = {group.id for group in user_groups} + + result = [] + for model in models: + if model.user_id == user_id: + result.append(model) + elif await AccessGrants.has_access( user_id=user_id, resource_type='model', resource_id=model.id, permission=permission, user_group_ids=user_group_ids, db=db, - ) - ] + ): + result.append(model) + return result def _has_permission(self, db, query, filter: dict, permission: str = 'read'): return AccessGrants.has_permission_filter( @@ -270,23 +274,22 @@ class ModelsTable: permission=permission, ) - def search_models( + async def search_models( self, user_id: str, filter: dict = {}, skip: int = 0, limit: int = 30, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> ModelListResponse: - with get_db_context(db) as db: - # Join GroupMember so we can order by group_id when requested - query = db.query(Model, User).outerjoin(User, User.id == Model.user_id) - query = query.filter(Model.base_model_id != None) + async with get_async_db_context(db) as db: + stmt = select(Model, User).outerjoin(User, User.id == Model.user_id) + stmt = stmt.filter(Model.base_model_id != None) if filter: query_key = filter.get('query') if query_key: - query = query.filter( + stmt = stmt.filter( or_( Model.name.ilike(f'%{query_key}%'), Model.base_model_id.ilike(f'%{query_key}%'), @@ -298,92 +301,95 @@ class ModelsTable: view_option = filter.get('view_option') if view_option == 'created': - query = query.filter(Model.user_id == user_id) + stmt = stmt.filter(Model.user_id == user_id) elif view_option == 'shared': - query = query.filter(Model.user_id != user_id) + stmt = stmt.filter(Model.user_id != user_id) # Apply access control filtering - query = self._has_permission( + stmt = self._has_permission( db, - query, + stmt, filter, permission='read', ) tag = filter.get('tag') if tag: - # TODO: This is a simple implementation and should be improved for performance - like_pattern = f'%"{tag.lower()}"%' # `"tag"` inside JSON array + like_pattern = f'%"{tag.lower()}"%' meta_text = func.lower(cast(Model.meta, String)) - - query = query.filter(meta_text.like(like_pattern)) + stmt = stmt.filter(meta_text.like(like_pattern)) order_by = filter.get('order_by') direction = filter.get('direction') if order_by == 'name': if direction == 'asc': - query = query.order_by(Model.name.asc()) + stmt = stmt.order_by(Model.name.asc()) else: - query = query.order_by(Model.name.desc()) + stmt = stmt.order_by(Model.name.desc()) elif order_by == 'created_at': if direction == 'asc': - query = query.order_by(Model.created_at.asc()) + stmt = stmt.order_by(Model.created_at.asc()) else: - query = query.order_by(Model.created_at.desc()) + stmt = stmt.order_by(Model.created_at.desc()) elif order_by == 'updated_at': if direction == 'asc': - query = query.order_by(Model.updated_at.asc()) + stmt = stmt.order_by(Model.updated_at.asc()) else: - query = query.order_by(Model.updated_at.desc()) + stmt = stmt.order_by(Model.updated_at.desc()) else: - query = query.order_by(Model.created_at.desc()) + stmt = stmt.order_by(Model.created_at.desc()) # Count BEFORE pagination - total = query.count() + count_result = await db.execute( + select(func.count()).select_from(stmt.subquery()) + ) + total = count_result.scalar() if skip: - query = query.offset(skip) + stmt = stmt.offset(skip) if limit: - query = query.limit(limit) + stmt = stmt.limit(limit) - items = query.all() + result = await db.execute(stmt) + items = result.all() model_ids = [model.id for model, _ in items] - grants_map = AccessGrants.get_grants_by_resources('model', model_ids, db=db) + grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db) models = [] for model, user in items: models.append( ModelUserResponse( - **self._to_model_model( + **(await self._to_model_model( model, access_grants=grants_map.get(model.id, []), db=db, - ).model_dump(), + )).model_dump(), user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None), ) ) return ModelListResponse(items=models, total=total) - def get_model_by_id(self, id: str, db: Optional[Session] = None) -> Optional[ModelModel]: + async def get_model_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[ModelModel]: try: - with get_db_context(db) as db: - model = db.get(Model, id) - return self._to_model_model(model, db=db) if model else None + async with get_async_db_context(db) as db: + model = await db.get(Model, id) + return await self._to_model_model(model, db=db) if model else None except Exception: return None - def get_models_by_ids(self, ids: list[str], db: Optional[Session] = None) -> list[ModelModel]: + async def get_models_by_ids(self, ids: list[str], db: Optional[AsyncSession] = None) -> list[ModelModel]: try: - with get_db_context(db) as db: - models = db.query(Model).filter(Model.id.in_(ids)).all() + async with get_async_db_context(db) as db: + result = await db.execute(select(Model).filter(Model.id.in_(ids))) + models = result.scalars().all() model_ids = [model.id for model in models] - grants_map = AccessGrants.get_grants_by_resources('model', model_ids, db=db) + grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db) return [ - self._to_model_model( + await self._to_model_model( model, access_grants=grants_map.get(model.id, []), db=db, @@ -393,82 +399,86 @@ class ModelsTable: except Exception: return [] - def toggle_model_by_id(self, id: str, db: Optional[Session] = None) -> Optional[ModelModel]: - with get_db_context(db) as db: + async def toggle_model_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[ModelModel]: + async with get_async_db_context(db) as db: try: - model = db.query(Model).filter_by(id=id).first() + result = await db.execute(select(Model).filter_by(id=id)) + model = result.scalars().first() if not model: return None model.is_active = not model.is_active model.updated_at = int(time.time()) - db.commit() - db.refresh(model) + await db.commit() + await db.refresh(model) - return self._to_model_model(model, db=db) + return await self._to_model_model(model, db=db) except Exception: return None - def update_model_by_id(self, id: str, model: ModelForm, db: Optional[Session] = None) -> Optional[ModelModel]: + async def update_model_by_id(self, id: str, model: ModelForm, db: Optional[AsyncSession] = None) -> Optional[ModelModel]: try: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: # update only the fields that are present in the model data = model.model_dump(exclude={'id', 'access_grants'}) data['updated_at'] = int(time.time()) - result = db.query(Model).filter_by(id=id).update(data) + await db.execute(update(Model).filter_by(id=id).values(**data)) - db.commit() + await db.commit() if model.access_grants is not None: - AccessGrants.set_access_grants('model', id, model.access_grants, db=db) + await AccessGrants.set_access_grants('model', id, model.access_grants, db=db) - return self.get_model_by_id(id, db=db) + return await self.get_model_by_id(id, db=db) except Exception as e: log.exception(f'Failed to update the model by id {id}: {e}') return None - def update_model_updated_at_by_id(self, id: str, db: Optional[Session] = None) -> Optional[ModelModel]: + async def update_model_updated_at_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[ModelModel]: try: - with get_db_context(db) as db: - result = db.query(Model).filter_by(id=id).first() - if not result: + async with get_async_db_context(db) as db: + result = await db.execute(select(Model).filter_by(id=id)) + model_obj = result.scalars().first() + if not model_obj: return None - result.updated_at = int(time.time()) - db.commit() - db.refresh(result) - return self._to_model_model(result, db=db) + model_obj.updated_at = int(time.time()) + await db.commit() + await db.refresh(model_obj) + return await self._to_model_model(model_obj, db=db) except Exception as e: log.exception(f'Failed to update the model updated_at by id {id}: {e}') return None - def delete_model_by_id(self, id: str, db: Optional[Session] = None) -> bool: + async def delete_model_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: - AccessGrants.revoke_all_access('model', id, db=db) - db.query(Model).filter_by(id=id).delete() - db.commit() + async with get_async_db_context(db) as db: + await AccessGrants.revoke_all_access('model', id, db=db) + await db.execute(delete(Model).filter_by(id=id)) + await db.commit() return True except Exception: return False - def delete_all_models(self, db: Optional[Session] = None) -> bool: + async def delete_all_models(self, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: - model_ids = [row[0] for row in db.query(Model.id).all()] + async with get_async_db_context(db) as db: + result = await db.execute(select(Model.id)) + model_ids = [row[0] for row in result.all()] for model_id in model_ids: - AccessGrants.revoke_all_access('model', model_id, db=db) - db.query(Model).delete() - db.commit() + await AccessGrants.revoke_all_access('model', model_id, db=db) + await db.execute(delete(Model)) + await db.commit() return True except Exception: return False - def sync_models(self, user_id: str, models: list[ModelModel], db: Optional[Session] = None) -> list[ModelModel]: + async def sync_models(self, user_id: str, models: list[ModelModel], db: Optional[AsyncSession] = None) -> list[ModelModel]: try: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: # Get existing models - existing_models = db.query(Model).all() + result = await db.execute(select(Model)) + existing_models = result.scalars().all() existing_ids = {model.id for model in existing_models} # Prepare a set of new model IDs @@ -477,12 +487,12 @@ class ModelsTable: # Update or insert models for model in models: if model.id in existing_ids: - db.query(Model).filter_by(id=model.id).update( - { + await db.execute( + update(Model).filter_by(id=model.id).values( **model.model_dump(exclude={'access_grants'}), - 'user_id': user_id, - 'updated_at': int(time.time()), - } + user_id=user_id, + updated_at=int(time.time()), + ) ) else: new_model = Model( @@ -493,21 +503,22 @@ class ModelsTable: } ) db.add(new_model) - AccessGrants.set_access_grants('model', model.id, model.access_grants, db=db) + await AccessGrants.set_access_grants('model', model.id, model.access_grants, db=db) # Remove models that are no longer present for model in existing_models: if model.id not in new_model_ids: - AccessGrants.revoke_all_access('model', model.id, db=db) - db.delete(model) + await AccessGrants.revoke_all_access('model', model.id, db=db) + await db.delete(model) - db.commit() + await db.commit() - all_models = db.query(Model).all() + result = await db.execute(select(Model)) + all_models = result.scalars().all() model_ids = [model.id for model in all_models] - grants_map = AccessGrants.get_grants_by_resources('model', model_ids, db=db) + grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db) return [ - self._to_model_model( + await self._to_model_model( model, access_grants=grants_map.get(model.id, []), db=db, diff --git a/backend/open_webui/models/notes.py b/backend/open_webui/models/notes.py index 34749f5f6c..b06465c7ad 100644 --- a/backend/open_webui/models/notes.py +++ b/backend/open_webui/models/notes.py @@ -4,8 +4,9 @@ import uuid from typing import Optional from functools import lru_cache -from sqlalchemy.orm import Session -from open_webui.internal.db import Base, get_db, get_db_context +from sqlalchemy import select, delete, update, or_, func, cast +from sqlalchemy.ext.asyncio import AsyncSession +from open_webui.internal.db import Base, get_async_db_context from open_webui.models.groups import Groups from open_webui.models.users import User, UserModel, Users, UserResponse from open_webui.models.access_grants import AccessGrantModel, AccessGrants @@ -13,7 +14,6 @@ from open_webui.models.access_grants import AccessGrantModel, AccessGrants from pydantic import BaseModel, ConfigDict, Field from sqlalchemy import BigInteger, Column, Text, JSON -from sqlalchemy import or_, func, cast #################### # Note DB Schema @@ -88,18 +88,18 @@ class NoteListResponse(BaseModel): class NoteTable: - def _get_access_grants(self, note_id: str, db: Optional[Session] = None) -> list[AccessGrantModel]: - return AccessGrants.get_grants_by_resource('note', note_id, db=db) + async def _get_access_grants(self, note_id: str, db: Optional[AsyncSession] = None) -> list[AccessGrantModel]: + return await AccessGrants.get_grants_by_resource('note', note_id, db=db) - def _to_note_model( + async def _to_note_model( self, note: Note, access_grants: Optional[list[AccessGrantModel]] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> NoteModel: note_data = NoteModel.model_validate(note).model_dump(exclude={'access_grants'}) note_data['access_grants'] = ( - access_grants if access_grants is not None else self._get_access_grants(note_data['id'], db=db) + access_grants if access_grants is not None else await self._get_access_grants(note_data['id'], db=db) ) return NoteModel.model_validate(note_data) @@ -113,8 +113,8 @@ class NoteTable: permission=permission, ) - def insert_new_note(self, user_id: str, form_data: NoteForm, db: Optional[Session] = None) -> Optional[NoteModel]: - with get_db_context(db) as db: + async def insert_new_note(self, user_id: str, form_data: NoteForm, db: Optional[AsyncSession] = None) -> Optional[NoteModel]: + async with get_async_db_context(db) as db: note = NoteModel( **{ 'id': str(uuid.uuid4()), @@ -129,38 +129,39 @@ class NoteTable: new_note = Note(**note.model_dump(exclude={'access_grants'})) db.add(new_note) - db.commit() - AccessGrants.set_access_grants('note', note.id, form_data.access_grants, db=db) - return self._to_note_model(new_note, db=db) + await db.commit() + await AccessGrants.set_access_grants('note', note.id, form_data.access_grants, db=db) + return await self._to_note_model(new_note, db=db) - def get_notes(self, skip: int = 0, limit: int = 50, db: Optional[Session] = None) -> list[NoteModel]: - with get_db_context(db) as db: - query = db.query(Note).order_by(Note.updated_at.desc()) + async def get_notes(self, skip: int = 0, limit: int = 50, db: Optional[AsyncSession] = None) -> list[NoteModel]: + async with get_async_db_context(db) as db: + stmt = select(Note).order_by(Note.updated_at.desc()) if skip is not None: - query = query.offset(skip) + stmt = stmt.offset(skip) if limit is not None: - query = query.limit(limit) - notes = query.all() + stmt = stmt.limit(limit) + result = await db.execute(stmt) + notes = result.scalars().all() note_ids = [note.id for note in notes] - grants_map = AccessGrants.get_grants_by_resources('note', note_ids, db=db) - return [self._to_note_model(note, access_grants=grants_map.get(note.id, []), db=db) for note in notes] + grants_map = await AccessGrants.get_grants_by_resources('note', note_ids, db=db) + return [await self._to_note_model(note, access_grants=grants_map.get(note.id, []), db=db) for note in notes] - def search_notes( + async def search_notes( self, user_id: str, filter: dict = {}, skip: int = 0, limit: int = 30, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> NoteListResponse: - with get_db_context(db) as db: - query = db.query(Note, User).outerjoin(User, User.id == Note.user_id) + async with get_async_db_context(db) as db: + stmt = select(Note, User).outerjoin(User, User.id == Note.user_id) if filter: query_key = filter.get('query') if query_key: # Normalize search by removing hyphens and spaces (e.g., "todo" matches "to-do" and "to do") normalized_query = query_key.replace('-', '').replace(' ', '') - query = query.filter( + stmt = stmt.filter( or_( func.replace(func.replace(Note.title, '-', ''), ' ', '').ilike(f'%{normalized_query}%'), func.replace( @@ -173,9 +174,9 @@ class NoteTable: view_option = filter.get('view_option') if view_option == 'created': - query = query.filter(Note.user_id == user_id) + stmt = stmt.filter(Note.user_id == user_id) elif view_option == 'shared': - query = query.filter(Note.user_id != user_id) + stmt = stmt.filter(Note.user_id != user_id) # Apply access control filtering if 'permission' in filter: @@ -183,9 +184,9 @@ class NoteTable: else: permission = 'write' - query = self._has_permission( + stmt = self._has_permission( db, - query, + stmt, filter, permission=permission, ) @@ -195,87 +196,95 @@ class NoteTable: if order_by == 'name': if direction == 'asc': - query = query.order_by(Note.title.asc()) + stmt = stmt.order_by(Note.title.asc()) else: - query = query.order_by(Note.title.desc()) + stmt = stmt.order_by(Note.title.desc()) elif order_by == 'created_at': if direction == 'asc': - query = query.order_by(Note.created_at.asc()) + stmt = stmt.order_by(Note.created_at.asc()) else: - query = query.order_by(Note.created_at.desc()) + stmt = stmt.order_by(Note.created_at.desc()) elif order_by == 'updated_at': if direction == 'asc': - query = query.order_by(Note.updated_at.asc()) + stmt = stmt.order_by(Note.updated_at.asc()) else: - query = query.order_by(Note.updated_at.desc()) + stmt = stmt.order_by(Note.updated_at.desc()) else: - query = query.order_by(Note.updated_at.desc()) + stmt = stmt.order_by(Note.updated_at.desc()) else: - query = query.order_by(Note.updated_at.desc()) + stmt = stmt.order_by(Note.updated_at.desc()) # Count BEFORE pagination - total = query.count() + count_result = await db.execute( + select(func.count()).select_from(stmt.subquery()) + ) + total = count_result.scalar() if skip: - query = query.offset(skip) + stmt = stmt.offset(skip) if limit: - query = query.limit(limit) + stmt = stmt.limit(limit) - items = query.all() + result = await db.execute(stmt) + items = result.all() note_ids = [note.id for note, _ in items] - grants_map = AccessGrants.get_grants_by_resources('note', note_ids, db=db) + grants_map = await AccessGrants.get_grants_by_resources('note', note_ids, db=db) notes = [] for note, user in items: notes.append( NoteUserResponse( - **self._to_note_model( + **(await self._to_note_model( note, access_grants=grants_map.get(note.id, []), db=db, - ).model_dump(), + )).model_dump(), user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None), ) ) return NoteListResponse(items=notes, total=total) - def get_notes_by_user_id( + async def get_notes_by_user_id( self, user_id: str, permission: str = 'read', skip: int = 0, limit: int = 50, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[NoteModel]: - with get_db_context(db) as db: - user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id, db=db)] + async with get_async_db_context(db) as db: + user_groups = await Groups.get_groups_by_member_id(user_id, db=db) + user_group_ids = [group.id for group in user_groups] - query = db.query(Note).order_by(Note.updated_at.desc()) - query = self._has_permission(db, query, {'user_id': user_id, 'group_ids': user_group_ids}, permission) + stmt = select(Note).order_by(Note.updated_at.desc()) + stmt = self._has_permission(db, stmt, {'user_id': user_id, 'group_ids': user_group_ids}, permission) if skip is not None: - query = query.offset(skip) + stmt = stmt.offset(skip) if limit is not None: - query = query.limit(limit) + stmt = stmt.limit(limit) - notes = query.all() + result = await db.execute(stmt) + notes = result.scalars().all() note_ids = [note.id for note in notes] - grants_map = AccessGrants.get_grants_by_resources('note', note_ids, db=db) - return [self._to_note_model(note, access_grants=grants_map.get(note.id, []), db=db) for note in notes] + grants_map = await AccessGrants.get_grants_by_resources('note', note_ids, db=db) + return [await self._to_note_model(note, access_grants=grants_map.get(note.id, []), db=db) for note in notes] - def get_note_by_id(self, id: str, db: Optional[Session] = None) -> Optional[NoteModel]: - with get_db_context(db) as db: - note = db.query(Note).filter(Note.id == id).first() - return self._to_note_model(note, db=db) if note else None + async def get_note_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[NoteModel]: + async with get_async_db_context(db) as db: + result = await db.execute(select(Note).filter(Note.id == id)) + note = result.scalars().first() + return await self._to_note_model(note, db=db) if note else None - def update_note_by_id( - self, id: str, form_data: NoteUpdateForm, db: Optional[Session] = None + async def update_note_by_id( + self, id: str, form_data: NoteUpdateForm, db: Optional[AsyncSession] = None ) -> Optional[NoteModel]: - with get_db_context(db) as db: - note = db.query(Note).filter(Note.id == id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(Note).filter(Note.id == id)) + note = result.scalars().first() if not note: return None @@ -289,19 +298,19 @@ class NoteTable: note.meta = {**note.meta, **form_data['meta']} if 'access_grants' in form_data: - AccessGrants.set_access_grants('note', id, form_data['access_grants'], db=db) + await AccessGrants.set_access_grants('note', id, form_data['access_grants'], db=db) note.updated_at = int(time.time_ns()) - db.commit() - return self._to_note_model(note, db=db) if note else None + await db.commit() + return await self._to_note_model(note, db=db) if note else None - def delete_note_by_id(self, id: str, db: Optional[Session] = None) -> bool: + async def delete_note_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: - AccessGrants.revoke_all_access('note', id, db=db) - db.query(Note).filter(Note.id == id).delete() - db.commit() + async with get_async_db_context(db) as db: + await AccessGrants.revoke_all_access('note', id, db=db) + await db.execute(delete(Note).filter(Note.id == id)) + await db.commit() return True except Exception: return False diff --git a/backend/open_webui/models/oauth_sessions.py b/backend/open_webui/models/oauth_sessions.py index 868216164a..c8ff569f27 100644 --- a/backend/open_webui/models/oauth_sessions.py +++ b/backend/open_webui/models/oauth_sessions.py @@ -8,8 +8,9 @@ import json from cryptography.fernet import Fernet -from sqlalchemy.orm import Session -from open_webui.internal.db import Base, get_db, get_db_context +from sqlalchemy import select, delete, update +from sqlalchemy.ext.asyncio import AsyncSession +from open_webui.internal.db import Base, get_async_db_context from open_webui.env import OAUTH_SESSION_TOKEN_ENCRYPTION_KEY from pydantic import BaseModel, ConfigDict @@ -103,16 +104,16 @@ class OAuthSessionTable: log.error(f'Error decrypting tokens: {type(e).__name__}: {e}') raise - def create_session( + async def create_session( self, user_id: str, provider: str, token: dict, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[OAuthSessionModel]: """Create a new OAuth session""" try: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: current_time = int(time.time()) id = str(uuid.uuid4()) @@ -129,91 +130,126 @@ class OAuthSessionTable: ) db.add(result) - db.commit() - db.refresh(result) + await db.commit() + await db.refresh(result) if result: - db.expunge(result) # Detach so dict swap is never flushed - result.token = token # Return decrypted token - return OAuthSessionModel.model_validate(result) + # Make a copy of the model data before closing session + model = OAuthSessionModel( + id=result.id, + user_id=result.user_id, + provider=result.provider, + token=token, # Return decrypted token + expires_at=result.expires_at, + created_at=result.created_at, + updated_at=result.updated_at, + ) + return model else: return None except Exception as e: log.error(f'Error creating OAuth session: {e}') return None - def get_session_by_id(self, session_id: str, db: Optional[Session] = None) -> Optional[OAuthSessionModel]: + async def get_session_by_id(self, session_id: str, db: Optional[AsyncSession] = None) -> Optional[OAuthSessionModel]: """Get OAuth session by ID""" try: - with get_db_context(db) as db: - session = db.query(OAuthSession).filter_by(id=session_id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(OAuthSession).filter_by(id=session_id)) + session = result.scalars().first() if session: - db.expunge(session) - session.token = self._decrypt_token(session.token) - return OAuthSessionModel.model_validate(session) + return OAuthSessionModel( + id=session.id, + user_id=session.user_id, + provider=session.provider, + token=self._decrypt_token(session.token), + expires_at=session.expires_at, + created_at=session.created_at, + updated_at=session.updated_at, + ) return None except Exception as e: log.error(f'Error getting OAuth session by ID: {e}') return None - def get_session_by_id_and_user_id( - self, session_id: str, user_id: str, db: Optional[Session] = None + async def get_session_by_id_and_user_id( + self, session_id: str, user_id: str, db: Optional[AsyncSession] = None ) -> Optional[OAuthSessionModel]: """Get OAuth session by ID and user ID""" try: - with get_db_context(db) as db: - session = db.query(OAuthSession).filter_by(id=session_id, user_id=user_id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(OAuthSession).filter_by(id=session_id, user_id=user_id)) + session = result.scalars().first() if session: - db.expunge(session) - session.token = self._decrypt_token(session.token) - return OAuthSessionModel.model_validate(session) + return OAuthSessionModel( + id=session.id, + user_id=session.user_id, + provider=session.provider, + token=self._decrypt_token(session.token), + expires_at=session.expires_at, + created_at=session.created_at, + updated_at=session.updated_at, + ) return None except Exception as e: log.error(f'Error getting OAuth session by ID: {e}') return None - def get_session_by_provider_and_user_id( - self, provider: str, user_id: str, db: Optional[Session] = None + async def get_session_by_provider_and_user_id( + self, provider: str, user_id: str, db: Optional[AsyncSession] = None ) -> Optional[OAuthSessionModel]: """Get OAuth session by provider and user ID""" try: - with get_db_context(db) as db: - session = ( - db.query(OAuthSession) + async with get_async_db_context(db) as db: + result = await db.execute( + select(OAuthSession) .filter_by(provider=provider, user_id=user_id) .order_by(OAuthSession.created_at.desc()) - .first() ) + session = result.scalars().first() if session: - db.expunge(session) - session.token = self._decrypt_token(session.token) - return OAuthSessionModel.model_validate(session) + return OAuthSessionModel( + id=session.id, + user_id=session.user_id, + provider=session.provider, + token=self._decrypt_token(session.token), + expires_at=session.expires_at, + created_at=session.created_at, + updated_at=session.updated_at, + ) return None except Exception as e: log.error(f'Error getting OAuth session by provider and user ID: {e}') return None - def get_sessions_by_user_id(self, user_id: str, db: Optional[Session] = None) -> List[OAuthSessionModel]: + async def get_sessions_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> List[OAuthSessionModel]: """Get all OAuth sessions for a user""" try: - with get_db_context(db) as db: - sessions = db.query(OAuthSession).filter_by(user_id=user_id).all() + async with get_async_db_context(db) as db: + result = await db.execute(select(OAuthSession).filter_by(user_id=user_id)) + sessions = result.scalars().all() results = [] for session in sessions: try: - db.expunge(session) - session.token = self._decrypt_token(session.token) - results.append(OAuthSessionModel.model_validate(session)) + results.append(OAuthSessionModel( + id=session.id, + user_id=session.user_id, + provider=session.provider, + token=self._decrypt_token(session.token), + expires_at=session.expires_at, + created_at=session.created_at, + updated_at=session.updated_at, + )) except Exception as e: log.warning( f'Skipping OAuth session {session.id} due to decryption failure, deleting corrupted session: {type(e).__name__}: {e}' ) - db.query(OAuthSession).filter_by(id=session.id).delete() - db.commit() + await db.execute(delete(OAuthSession).filter_by(id=session.id)) + await db.commit() return results @@ -221,62 +257,69 @@ class OAuthSessionTable: log.error(f'Error getting OAuth sessions by user ID: {e}') return [] - def update_session_by_id( - self, session_id: str, token: dict, db: Optional[Session] = None + async def update_session_by_id( + self, session_id: str, token: dict, db: Optional[AsyncSession] = None ) -> Optional[OAuthSessionModel]: """Update OAuth session tokens""" try: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: current_time = int(time.time()) - db.query(OAuthSession).filter_by(id=session_id).update( - { - 'token': self._encrypt_token(token), - 'expires_at': token.get('expires_at'), - 'updated_at': current_time, - } + await db.execute( + update(OAuthSession).filter_by(id=session_id).values( + token=self._encrypt_token(token), + expires_at=token.get('expires_at'), + updated_at=current_time, + ) ) - db.commit() - session = db.query(OAuthSession).filter_by(id=session_id).first() + await db.commit() + result = await db.execute(select(OAuthSession).filter_by(id=session_id)) + session = result.scalars().first() if session: - db.expunge(session) - session.token = self._decrypt_token(session.token) - return OAuthSessionModel.model_validate(session) + return OAuthSessionModel( + id=session.id, + user_id=session.user_id, + provider=session.provider, + token=self._decrypt_token(session.token), + expires_at=session.expires_at, + created_at=session.created_at, + updated_at=session.updated_at, + ) return None except Exception as e: log.error(f'Error updating OAuth session tokens: {e}') return None - def delete_session_by_id(self, session_id: str, db: Optional[Session] = None) -> bool: + async def delete_session_by_id(self, session_id: str, db: Optional[AsyncSession] = None) -> bool: """Delete an OAuth session""" try: - with get_db_context(db) as db: - result = db.query(OAuthSession).filter_by(id=session_id).delete() - db.commit() - return result > 0 + async with get_async_db_context(db) as db: + result = await db.execute(delete(OAuthSession).filter_by(id=session_id)) + await db.commit() + return result.rowcount > 0 except Exception as e: log.error(f'Error deleting OAuth session: {e}') return False - def delete_sessions_by_user_id(self, user_id: str, db: Optional[Session] = None) -> bool: + async def delete_sessions_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> bool: """Delete all OAuth sessions for a user""" try: - with get_db_context(db) as db: - result = db.query(OAuthSession).filter_by(user_id=user_id).delete() - db.commit() + async with get_async_db_context(db) as db: + await db.execute(delete(OAuthSession).filter_by(user_id=user_id)) + await db.commit() return True except Exception as e: log.error(f'Error deleting OAuth sessions by user ID: {e}') return False - def delete_sessions_by_provider(self, provider: str, db: Optional[Session] = None) -> bool: + async def delete_sessions_by_provider(self, provider: str, db: Optional[AsyncSession] = None) -> bool: """Delete all OAuth sessions for a provider""" try: - with get_db_context(db) as db: - db.query(OAuthSession).filter_by(provider=provider).delete() - db.commit() + async with get_async_db_context(db) as db: + await db.execute(delete(OAuthSession).filter_by(provider=provider)) + await db.commit() return True except Exception as e: log.error(f'Error deleting OAuth sessions by provider {provider}: {e}') diff --git a/backend/open_webui/models/prompt_history.py b/backend/open_webui/models/prompt_history.py index d42b4bfa24..5d0f4a65b2 100644 --- a/backend/open_webui/models/prompt_history.py +++ b/backend/open_webui/models/prompt_history.py @@ -6,8 +6,9 @@ from typing import Optional import json import difflib -from sqlalchemy.orm import Session -from open_webui.internal.db import Base, get_db_context +from sqlalchemy import select, delete, func +from sqlalchemy.ext.asyncio import AsyncSession +from open_webui.internal.db import Base, get_async_db_context from open_webui.models.users import Users, UserResponse from pydantic import BaseModel, ConfigDict @@ -49,17 +50,17 @@ class PromptHistoryResponse(PromptHistoryModel): class PromptHistoryTable: - def create_history_entry( + async def create_history_entry( self, prompt_id: str, snapshot: dict, user_id: str, parent_id: Optional[str] = None, commit_message: Optional[str] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[PromptHistoryModel]: """Create a new history entry (commit) for a prompt.""" - with get_db_context(db) as db: + async with get_async_db_context(db) as db: history = PromptHistory( id=str(uuid.uuid4()), prompt_id=prompt_id, @@ -70,31 +71,31 @@ class PromptHistoryTable: created_at=int(time.time()), ) db.add(history) - db.commit() - db.refresh(history) + await db.commit() + await db.refresh(history) return PromptHistoryModel.model_validate(history) - def get_history_by_prompt_id( + async def get_history_by_prompt_id( self, prompt_id: str, limit: int = 50, offset: int = 0, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[PromptHistoryResponse]: """Get all history entries for a prompt, ordered by created_at desc.""" - with get_db_context(db) as db: - entries = ( - db.query(PromptHistory) + async with get_async_db_context(db) as db: + result = await db.execute( + select(PromptHistory) .filter(PromptHistory.prompt_id == prompt_id) .order_by(PromptHistory.created_at.desc()) .offset(offset) .limit(limit) - .all() ) + entries = result.scalars().all() # Get user info for each entry user_ids = list(set(e.user_id for e in entries)) - users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else [] + users = await Users.get_users_by_user_ids(user_ids, db=db) if user_ids else [] users_dict = {user.id: user for user in users} return [ @@ -105,54 +106,61 @@ class PromptHistoryTable: for entry in entries ] - def get_history_entry_by_id( + async def get_history_entry_by_id( self, history_id: str, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[PromptHistoryModel]: """Get a specific history entry by ID.""" - with get_db_context(db) as db: - entry = db.query(PromptHistory).filter(PromptHistory.id == history_id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(PromptHistory).filter(PromptHistory.id == history_id)) + entry = result.scalars().first() if entry: return PromptHistoryModel.model_validate(entry) return None - def get_latest_history_entry( + async def get_latest_history_entry( self, prompt_id: str, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[PromptHistoryModel]: """Get the most recent history entry for a prompt.""" - with get_db_context(db) as db: - entry = ( - db.query(PromptHistory) + async with get_async_db_context(db) as db: + result = await db.execute( + select(PromptHistory) .filter(PromptHistory.prompt_id == prompt_id) .order_by(PromptHistory.created_at.desc()) - .first() + .limit(1) ) + entry = result.scalars().first() if entry: return PromptHistoryModel.model_validate(entry) return None - def get_history_count( + async def get_history_count( self, prompt_id: str, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> int: """Get the number of history entries for a prompt.""" - with get_db_context(db) as db: - return db.query(PromptHistory).filter(PromptHistory.prompt_id == prompt_id).count() + async with get_async_db_context(db) as db: + result = await db.execute( + select(func.count()).select_from(PromptHistory).filter(PromptHistory.prompt_id == prompt_id) + ) + return result.scalar() - def compute_diff( + async def compute_diff( self, from_id: str, to_id: str, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[dict]: """Compute diff between two history entries.""" - with get_db_context(db) as db: - from_entry = db.query(PromptHistory).filter(PromptHistory.id == from_id).first() - to_entry = db.query(PromptHistory).filter(PromptHistory.id == to_id).first() + async with get_async_db_context(db) as db: + result_from = await db.execute(select(PromptHistory).filter(PromptHistory.id == from_id)) + from_entry = result_from.scalars().first() + result_to = await db.execute(select(PromptHistory).filter(PromptHistory.id == to_id)) + to_entry = result_to.scalars().first() if not from_entry or not to_entry: return None @@ -183,37 +191,39 @@ class PromptHistoryTable: 'name_changed': from_snapshot.get('name') != to_snapshot.get('name'), } - def delete_history_by_prompt_id( + async def delete_history_by_prompt_id( self, prompt_id: str, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> bool: """Delete all history entries for a prompt.""" - with get_db_context(db) as db: - db.query(PromptHistory).filter(PromptHistory.prompt_id == prompt_id).delete() - db.commit() + async with get_async_db_context(db) as db: + await db.execute(delete(PromptHistory).filter(PromptHistory.prompt_id == prompt_id)) + await db.commit() return True - def delete_history_entry( + async def delete_history_entry( self, history_id: str, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> bool: """Delete a history entry and reparent its children to grandparent.""" - with get_db_context(db) as db: - entry = db.query(PromptHistory).filter_by(id=history_id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(PromptHistory).filter_by(id=history_id)) + entry = result.scalars().first() if not entry: return False # Find children that reference this entry as parent - children = db.query(PromptHistory).filter_by(parent_id=history_id).all() + children_result = await db.execute(select(PromptHistory).filter_by(parent_id=history_id)) + children = children_result.scalars().all() # Reparent children to grandparent for child in children: child.parent_id = entry.parent_id - db.delete(entry) - db.commit() + await db.delete(entry) + await db.commit() return True diff --git a/backend/open_webui/models/prompts.py b/backend/open_webui/models/prompts.py index bb77f32f31..7250d1901e 100644 --- a/backend/open_webui/models/prompts.py +++ b/backend/open_webui/models/prompts.py @@ -2,16 +2,17 @@ import time import uuid from typing import Optional -from sqlalchemy.orm import Session -from open_webui.internal.db import Base, JSONField, get_db, get_db_context +from sqlalchemy import select, delete, update, or_, func, cast, String +from sqlalchemy.ext.asyncio import AsyncSession +from open_webui.internal.db import Base, JSONField, get_async_db_context from open_webui.models.groups import Groups -from open_webui.models.users import Users, UserResponse +from open_webui.models.users import Users, User, UserModel, UserResponse from open_webui.models.prompt_history import PromptHistories from open_webui.models.access_grants import AccessGrantModel, AccessGrants from pydantic import BaseModel, ConfigDict, Field -from sqlalchemy import BigInteger, Boolean, Column, String, Text, JSON, or_, func, cast +from sqlalchemy import BigInteger, Boolean, Column, Text, JSON #################### # Prompts DB Schema @@ -92,23 +93,23 @@ class PromptForm(BaseModel): class PromptsTable: - def _get_access_grants(self, prompt_id: str, db: Optional[Session] = None) -> list[AccessGrantModel]: - return AccessGrants.get_grants_by_resource('prompt', prompt_id, db=db) + async def _get_access_grants(self, prompt_id: str, db: Optional[AsyncSession] = None) -> list[AccessGrantModel]: + return await AccessGrants.get_grants_by_resource('prompt', prompt_id, db=db) - def _to_prompt_model( + async def _to_prompt_model( self, prompt: Prompt, access_grants: Optional[list[AccessGrantModel]] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> PromptModel: prompt_data = PromptModel.model_validate(prompt).model_dump(exclude={'access_grants'}) prompt_data['access_grants'] = ( - access_grants if access_grants is not None else self._get_access_grants(prompt_data['id'], db=db) + access_grants if access_grants is not None else await self._get_access_grants(prompt_data['id'], db=db) ) return PromptModel.model_validate(prompt_data) - def insert_new_prompt( - self, user_id: str, form_data: PromptForm, db: Optional[Session] = None + async def insert_new_prompt( + self, user_id: str, form_data: PromptForm, db: Optional[AsyncSession] = None ) -> Optional[PromptModel]: now = int(time.time()) prompt_id = str(uuid.uuid4()) @@ -129,15 +130,15 @@ class PromptsTable: ) try: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: result = Prompt(**prompt.model_dump(exclude={'access_grants'})) db.add(result) - db.commit() - db.refresh(result) - AccessGrants.set_access_grants('prompt', prompt_id, form_data.access_grants, db=db) + await db.commit() + await db.refresh(result) + await AccessGrants.set_access_grants('prompt', prompt_id, form_data.access_grants, db=db) if result: - current_access_grants = self._get_access_grants(prompt_id, db=db) + current_access_grants = await self._get_access_grants(prompt_id, db=db) snapshot = { 'name': form_data.name, 'content': form_data.content, @@ -148,7 +149,7 @@ class PromptsTable: 'access_grants': [grant.model_dump() for grant in current_access_grants], } - history_entry = PromptHistories.create_history_entry( + history_entry = await PromptHistories.create_history_entry( prompt_id=prompt_id, snapshot=snapshot, user_id=user_id, @@ -160,46 +161,51 @@ class PromptsTable: # Set the initial version as the production version if history_entry: result.version_id = history_entry.id - db.commit() - db.refresh(result) + await db.commit() + await db.refresh(result) - return self._to_prompt_model(result, db=db) + return await self._to_prompt_model(result, db=db) else: return None except Exception: return None - def get_prompt_by_id(self, prompt_id: str, db: Optional[Session] = None) -> Optional[PromptModel]: + async def get_prompt_by_id(self, prompt_id: str, db: Optional[AsyncSession] = None) -> Optional[PromptModel]: """Get prompt by UUID.""" try: - with get_db_context(db) as db: - prompt = db.query(Prompt).filter_by(id=prompt_id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(Prompt).filter_by(id=prompt_id)) + prompt = result.scalars().first() if prompt: - return self._to_prompt_model(prompt, db=db) + return await self._to_prompt_model(prompt, db=db) return None except Exception: return None - def get_prompt_by_command(self, command: str, db: Optional[Session] = None) -> Optional[PromptModel]: + async def get_prompt_by_command(self, command: str, db: Optional[AsyncSession] = None) -> Optional[PromptModel]: try: - with get_db_context(db) as db: - prompt = db.query(Prompt).filter_by(command=command).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(Prompt).filter_by(command=command)) + prompt = result.scalars().first() if prompt: - return self._to_prompt_model(prompt, db=db) + return await self._to_prompt_model(prompt, db=db) return None except Exception: return None - def get_prompts(self, db: Optional[Session] = None) -> list[PromptUserResponse]: - with get_db_context(db) as db: - all_prompts = db.query(Prompt).filter(Prompt.is_active == True).order_by(Prompt.updated_at.desc()).all() + async def get_prompts(self, db: Optional[AsyncSession] = None) -> list[PromptUserResponse]: + async with get_async_db_context(db) as db: + result = await db.execute( + select(Prompt).filter(Prompt.is_active == True).order_by(Prompt.updated_at.desc()) + ) + all_prompts = result.scalars().all() user_ids = list(set(prompt.user_id for prompt in all_prompts)) prompt_ids = [prompt.id for prompt in all_prompts] - users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else [] + users = await Users.get_users_by_user_ids(user_ids, db=db) if user_ids else [] users_dict = {user.id: user for user in users} - grants_map = AccessGrants.get_grants_by_resources('prompt', prompt_ids, db=db) + grants_map = await AccessGrants.get_grants_by_resources('prompt', prompt_ids, db=db) prompts = [] for prompt in all_prompts: @@ -207,11 +213,11 @@ class PromptsTable: prompts.append( PromptUserResponse.model_validate( { - **self._to_prompt_model( + **(await self._to_prompt_model( prompt, access_grants=grants_map.get(prompt.id, []), db=db, - ).model_dump(), + )).model_dump(), 'user': user.model_dump() if user else None, } ) @@ -219,44 +225,44 @@ class PromptsTable: return prompts - def get_prompts_by_user_id( - self, user_id: str, permission: str = 'write', db: Optional[Session] = None + async def get_prompts_by_user_id( + self, user_id: str, permission: str = 'write', db: Optional[AsyncSession] = None ) -> list[PromptUserResponse]: - prompts = self.get_prompts(db=db) - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id, db=db)} + prompts = await self.get_prompts(db=db) + user_groups = await Groups.get_groups_by_member_id(user_id, db=db) + user_group_ids = {group.id for group in user_groups} - return [ - prompt - for prompt in prompts - if prompt.user_id == user_id - or AccessGrants.has_access( + result = [] + for prompt in prompts: + if prompt.user_id == user_id: + result.append(prompt) + elif await AccessGrants.has_access( user_id=user_id, resource_type='prompt', resource_id=prompt.id, permission=permission, user_group_ids=user_group_ids, db=db, - ) - ] + ): + result.append(prompt) + return result - def search_prompts( + async def search_prompts( self, user_id: str, filter: dict = {}, skip: int = 0, limit: int = 30, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> PromptListResponse: - with get_db_context(db) as db: - from open_webui.models.users import User, UserModel - + async with get_async_db_context(db) as db: # Join with User table for user filtering and sorting - query = db.query(Prompt, User).outerjoin(User, User.id == Prompt.user_id) + stmt = select(Prompt, User).outerjoin(User, User.id == Prompt.user_id) if filter: query_key = filter.get('query') if query_key: - query = query.filter( + stmt = stmt.filter( or_( Prompt.name.ilike(f'%{query_key}%'), Prompt.command.ilike(f'%{query_key}%'), @@ -268,14 +274,14 @@ class PromptsTable: view_option = filter.get('view_option') if view_option == 'created': - query = query.filter(Prompt.user_id == user_id) + stmt = stmt.filter(Prompt.user_id == user_id) elif view_option == 'shared': - query = query.filter(Prompt.user_id != user_id) + stmt = stmt.filter(Prompt.user_id != user_id) # Apply access grant filtering - query = AccessGrants.has_permission_filter( + stmt = AccessGrants.has_permission_filter( db=db, - query=query, + query=stmt, DocumentModel=Prompt, filter=filter, resource_type='prompt', @@ -287,75 +293,80 @@ class PromptsTable: # Search for tag in JSON array field like_pattern = f'%"{tag.lower()}"%' tags_text = func.lower(cast(Prompt.tags, String)) - query = query.filter(tags_text.like(like_pattern)) + stmt = stmt.filter(tags_text.like(like_pattern)) order_by = filter.get('order_by') direction = filter.get('direction') if order_by == 'name': if direction == 'asc': - query = query.order_by(Prompt.name.asc()) + stmt = stmt.order_by(Prompt.name.asc()) else: - query = query.order_by(Prompt.name.desc()) + stmt = stmt.order_by(Prompt.name.desc()) elif order_by == 'created_at': if direction == 'asc': - query = query.order_by(Prompt.created_at.asc()) + stmt = stmt.order_by(Prompt.created_at.asc()) else: - query = query.order_by(Prompt.created_at.desc()) + stmt = stmt.order_by(Prompt.created_at.desc()) elif order_by == 'updated_at': if direction == 'asc': - query = query.order_by(Prompt.updated_at.asc()) + stmt = stmt.order_by(Prompt.updated_at.asc()) else: - query = query.order_by(Prompt.updated_at.desc()) + stmt = stmt.order_by(Prompt.updated_at.desc()) else: - query = query.order_by(Prompt.updated_at.desc()) + stmt = stmt.order_by(Prompt.updated_at.desc()) else: - query = query.order_by(Prompt.updated_at.desc()) + stmt = stmt.order_by(Prompt.updated_at.desc()) # Count BEFORE pagination - total = query.count() + count_result = await db.execute( + select(func.count()).select_from(stmt.subquery()) + ) + total = count_result.scalar() if skip: - query = query.offset(skip) + stmt = stmt.offset(skip) if limit: - query = query.limit(limit) + stmt = stmt.limit(limit) - items = query.all() + result = await db.execute(stmt) + items = result.all() prompt_ids = [prompt.id for prompt, _ in items] - grants_map = AccessGrants.get_grants_by_resources('prompt', prompt_ids, db=db) + grants_map = await AccessGrants.get_grants_by_resources('prompt', prompt_ids, db=db) prompts = [] for prompt, user in items: prompts.append( PromptUserResponse( - **self._to_prompt_model( + **(await self._to_prompt_model( prompt, access_grants=grants_map.get(prompt.id, []), db=db, - ).model_dump(), + )).model_dump(), user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None), ) ) return PromptListResponse(items=prompts, total=total) - def update_prompt_by_command( + async def update_prompt_by_command( self, command: str, form_data: PromptForm, user_id: str, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[PromptModel]: try: - with get_db_context(db) as db: - prompt = db.query(Prompt).filter_by(command=command).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(Prompt).filter_by(command=command)) + prompt = result.scalars().first() if not prompt: return None - latest_history = PromptHistories.get_latest_history_entry(prompt.id, db=db) + latest_history = await PromptHistories.get_latest_history_entry(prompt.id, db=db) parent_id = latest_history.id if latest_history else None - current_access_grants = self._get_access_grants(prompt.id, db=db) + current_access_grants = await self._get_access_grants(prompt.id, db=db) # Check if content changed to decide on history creation content_changed = ( @@ -371,10 +382,10 @@ class PromptsTable: prompt.meta = form_data.meta or prompt.meta prompt.updated_at = int(time.time()) if form_data.access_grants is not None: - AccessGrants.set_access_grants('prompt', prompt.id, form_data.access_grants, db=db) - current_access_grants = self._get_access_grants(prompt.id, db=db) + await AccessGrants.set_access_grants('prompt', prompt.id, form_data.access_grants, db=db) + current_access_grants = await self._get_access_grants(prompt.id, db=db) - db.commit() + await db.commit() # Create history entry only if content changed if content_changed: @@ -387,7 +398,7 @@ class PromptsTable: 'access_grants': [grant.model_dump() for grant in current_access_grants], } - history_entry = PromptHistories.create_history_entry( + history_entry = await PromptHistories.create_history_entry( prompt_id=prompt.id, snapshot=snapshot, user_id=user_id, @@ -399,28 +410,29 @@ class PromptsTable: # Set as production if flag is True (default) if form_data.is_production and history_entry: prompt.version_id = history_entry.id - db.commit() + await db.commit() - return self._to_prompt_model(prompt, db=db) + return await self._to_prompt_model(prompt, db=db) except Exception: return None - def update_prompt_by_id( + async def update_prompt_by_id( self, prompt_id: str, form_data: PromptForm, user_id: str, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[PromptModel]: try: - with get_db_context(db) as db: - prompt = db.query(Prompt).filter_by(id=prompt_id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(Prompt).filter_by(id=prompt_id)) + prompt = result.scalars().first() if not prompt: return None - latest_history = PromptHistories.get_latest_history_entry(prompt.id, db=db) + latest_history = await PromptHistories.get_latest_history_entry(prompt.id, db=db) parent_id = latest_history.id if latest_history else None - current_access_grants = self._get_access_grants(prompt.id, db=db) + current_access_grants = await self._get_access_grants(prompt.id, db=db) # Check if content changed to decide on history creation content_changed = ( @@ -442,12 +454,12 @@ class PromptsTable: prompt.tags = form_data.tags if form_data.access_grants is not None: - AccessGrants.set_access_grants('prompt', prompt.id, form_data.access_grants, db=db) - current_access_grants = self._get_access_grants(prompt.id, db=db) + await AccessGrants.set_access_grants('prompt', prompt.id, form_data.access_grants, db=db) + current_access_grants = await self._get_access_grants(prompt.id, db=db) prompt.updated_at = int(time.time()) - db.commit() + await db.commit() # Create history entry only if content changed if content_changed: @@ -461,7 +473,7 @@ class PromptsTable: 'access_grants': [grant.model_dump() for grant in current_access_grants], } - history_entry = PromptHistories.create_history_entry( + history_entry = await PromptHistories.create_history_entry( prompt_id=prompt.id, snapshot=snapshot, user_id=user_id, @@ -473,24 +485,25 @@ class PromptsTable: # Set as production if flag is True (default) if form_data.is_production and history_entry: prompt.version_id = history_entry.id - db.commit() + await db.commit() - return self._to_prompt_model(prompt, db=db) + return await self._to_prompt_model(prompt, db=db) except Exception: return None - def update_prompt_metadata( + async def update_prompt_metadata( self, prompt_id: str, name: str, command: str, tags: Optional[list[str]] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[PromptModel]: """Update only name, command, and tags (no history created).""" try: - with get_db_context(db) as db: - prompt = db.query(Prompt).filter_by(id=prompt_id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(Prompt).filter_by(id=prompt_id)) + prompt = result.scalars().first() if not prompt: return None @@ -501,26 +514,27 @@ class PromptsTable: prompt.tags = tags prompt.updated_at = int(time.time()) - db.commit() + await db.commit() - return self._to_prompt_model(prompt, db=db) + return await self._to_prompt_model(prompt, db=db) except Exception: return None - def update_prompt_version( + async def update_prompt_version( self, prompt_id: str, version_id: str, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[PromptModel]: """Set the active version of a prompt and restore content from that version's snapshot.""" try: - with get_db_context(db) as db: - prompt = db.query(Prompt).filter_by(id=prompt_id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(Prompt).filter_by(id=prompt_id)) + prompt = result.scalars().first() if not prompt: return None - history_entry = PromptHistories.get_history_entry_by_id(version_id, db=db) + history_entry = await PromptHistories.get_history_entry_by_id(version_id, db=db) if not history_entry: return None @@ -537,63 +551,67 @@ class PromptsTable: prompt.version_id = version_id prompt.updated_at = int(time.time()) - db.commit() + await db.commit() - return self._to_prompt_model(prompt, db=db) + return await self._to_prompt_model(prompt, db=db) except Exception: return None - def toggle_prompt_active(self, prompt_id: str, db: Optional[Session] = None) -> Optional[PromptModel]: + async def toggle_prompt_active(self, prompt_id: str, db: Optional[AsyncSession] = None) -> Optional[PromptModel]: """Toggle the is_active flag on a prompt.""" try: - with get_db_context(db) as db: - prompt = db.query(Prompt).filter_by(id=prompt_id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(Prompt).filter_by(id=prompt_id)) + prompt = result.scalars().first() if prompt: prompt.is_active = not prompt.is_active prompt.updated_at = int(time.time()) - db.commit() - db.refresh(prompt) - return self._to_prompt_model(prompt, db=db) + await db.commit() + await db.refresh(prompt) + return await self._to_prompt_model(prompt, db=db) return None except Exception: return None - def delete_prompt_by_command(self, command: str, db: Optional[Session] = None) -> bool: + async def delete_prompt_by_command(self, command: str, db: Optional[AsyncSession] = None) -> bool: """Permanently delete a prompt and its history.""" try: - with get_db_context(db) as db: - prompt = db.query(Prompt).filter_by(command=command).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(Prompt).filter_by(command=command)) + prompt = result.scalars().first() if prompt: - PromptHistories.delete_history_by_prompt_id(prompt.id, db=db) - AccessGrants.revoke_all_access('prompt', prompt.id, db=db) + await PromptHistories.delete_history_by_prompt_id(prompt.id, db=db) + await AccessGrants.revoke_all_access('prompt', prompt.id, db=db) - db.delete(prompt) - db.commit() + await db.delete(prompt) + await db.commit() return True return False except Exception: return False - def delete_prompt_by_id(self, prompt_id: str, db: Optional[Session] = None) -> bool: + async def delete_prompt_by_id(self, prompt_id: str, db: Optional[AsyncSession] = None) -> bool: """Permanently delete a prompt and its history.""" try: - with get_db_context(db) as db: - prompt = db.query(Prompt).filter_by(id=prompt_id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(Prompt).filter_by(id=prompt_id)) + prompt = result.scalars().first() if prompt: - PromptHistories.delete_history_by_prompt_id(prompt.id, db=db) - AccessGrants.revoke_all_access('prompt', prompt.id, db=db) + await PromptHistories.delete_history_by_prompt_id(prompt.id, db=db) + await AccessGrants.revoke_all_access('prompt', prompt.id, db=db) - db.delete(prompt) - db.commit() + await db.delete(prompt) + await db.commit() return True return False except Exception: return False - def get_tags(self, db: Optional[Session] = None) -> list[str]: + async def get_tags(self, db: Optional[AsyncSession] = None) -> list[str]: try: - with get_db_context(db) as db: - prompts = db.query(Prompt).filter_by(is_active=True).all() + async with get_async_db_context(db) as db: + result = await db.execute(select(Prompt).filter_by(is_active=True)) + prompts = result.scalars().all() tags = set() for prompt in prompts: if prompt.tags: diff --git a/backend/open_webui/models/skills.py b/backend/open_webui/models/skills.py index cdf8ecaea4..55ba204135 100644 --- a/backend/open_webui/models/skills.py +++ b/backend/open_webui/models/skills.py @@ -2,14 +2,15 @@ import logging import time from typing import Optional -from sqlalchemy.orm import Session -from open_webui.internal.db import Base, get_db, get_db_context -from open_webui.models.users import Users, UserResponse +from sqlalchemy import select, delete, update, or_ +from sqlalchemy.ext.asyncio import AsyncSession +from open_webui.internal.db import Base, get_async_db_context +from open_webui.models.users import Users, User, UserModel, UserResponse from open_webui.models.groups import Groups from open_webui.models.access_grants import AccessGrantModel, AccessGrants from pydantic import BaseModel, ConfigDict, Field -from sqlalchemy import JSON, BigInteger, Boolean, Column, String, Text, or_ +from sqlalchemy import JSON, BigInteger, Boolean, Column, String, Text, func log = logging.getLogger(__name__) @@ -105,28 +106,28 @@ class SkillAccessListResponse(BaseModel): class SkillsTable: - def _get_access_grants(self, skill_id: str, db: Optional[Session] = None) -> list[AccessGrantModel]: - return AccessGrants.get_grants_by_resource('skill', skill_id, db=db) + async def _get_access_grants(self, skill_id: str, db: Optional[AsyncSession] = None) -> list[AccessGrantModel]: + return await AccessGrants.get_grants_by_resource('skill', skill_id, db=db) - def _to_skill_model( + async def _to_skill_model( self, skill: Skill, access_grants: Optional[list[AccessGrantModel]] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> SkillModel: skill_data = SkillModel.model_validate(skill).model_dump(exclude={'access_grants'}) skill_data['access_grants'] = ( - access_grants if access_grants is not None else self._get_access_grants(skill_data['id'], db=db) + access_grants if access_grants is not None else await self._get_access_grants(skill_data['id'], db=db) ) return SkillModel.model_validate(skill_data) - def insert_new_skill( + async def insert_new_skill( self, user_id: str, form_data: SkillForm, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[SkillModel]: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: try: result = Skill( **{ @@ -137,43 +138,45 @@ class SkillsTable: } ) db.add(result) - db.commit() - db.refresh(result) - AccessGrants.set_access_grants('skill', result.id, form_data.access_grants, db=db) + await db.commit() + await db.refresh(result) + await AccessGrants.set_access_grants('skill', result.id, form_data.access_grants, db=db) if result: - return self._to_skill_model(result, db=db) + return await self._to_skill_model(result, db=db) else: return None except Exception as e: log.exception(f'Error creating a new skill: {e}') return None - def get_skill_by_id(self, id: str, db: Optional[Session] = None) -> Optional[SkillModel]: + async def get_skill_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[SkillModel]: try: - with get_db_context(db) as db: - skill = db.get(Skill, id) - return self._to_skill_model(skill, db=db) if skill else None + async with get_async_db_context(db) as db: + skill = await db.get(Skill, id) + return await self._to_skill_model(skill, db=db) if skill else None except Exception: return None - def get_skill_by_name(self, name: str, db: Optional[Session] = None) -> Optional[SkillModel]: + async def get_skill_by_name(self, name: str, db: Optional[AsyncSession] = None) -> Optional[SkillModel]: try: - with get_db_context(db) as db: - skill = db.query(Skill).filter_by(name=name).first() - return self._to_skill_model(skill, db=db) if skill else None + async with get_async_db_context(db) as db: + result = await db.execute(select(Skill).filter_by(name=name)) + skill = result.scalars().first() + return await self._to_skill_model(skill, db=db) if skill else None except Exception: return None - def get_skills(self, db: Optional[Session] = None) -> list[SkillUserModel]: - with get_db_context(db) as db: - all_skills = db.query(Skill).order_by(Skill.updated_at.desc()).all() + async def get_skills(self, db: Optional[AsyncSession] = None) -> list[SkillUserModel]: + async with get_async_db_context(db) as db: + result = await db.execute(select(Skill).order_by(Skill.updated_at.desc())) + all_skills = result.scalars().all() user_ids = list(set(skill.user_id for skill in all_skills)) skill_ids = [skill.id for skill in all_skills] - users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else [] + users = await Users.get_users_by_user_ids(user_ids, db=db) if user_ids else [] users_dict = {user.id: user for user in users} - grants_map = AccessGrants.get_grants_by_resources('skill', skill_ids, db=db) + grants_map = await AccessGrants.get_grants_by_resources('skill', skill_ids, db=db) skills = [] for skill in all_skills: @@ -181,56 +184,56 @@ class SkillsTable: skills.append( SkillUserModel.model_validate( { - **self._to_skill_model( + **(await self._to_skill_model( skill, access_grants=grants_map.get(skill.id, []), db=db, - ).model_dump(), + )).model_dump(), 'user': user.model_dump() if user else None, } ) ) return skills - def get_skills_by_user_id( - self, user_id: str, permission: str = 'write', db: Optional[Session] = None + async def get_skills_by_user_id( + self, user_id: str, permission: str = 'write', db: Optional[AsyncSession] = None ) -> list[SkillUserModel]: - skills = self.get_skills(db=db) - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id, db=db)} + skills = await self.get_skills(db=db) + user_groups = await Groups.get_groups_by_member_id(user_id, db=db) + user_group_ids = {group.id for group in user_groups} - return [ - skill - for skill in skills - if skill.user_id == user_id - or AccessGrants.has_access( + result = [] + for skill in skills: + if skill.user_id == user_id: + result.append(skill) + elif await AccessGrants.has_access( user_id=user_id, resource_type='skill', resource_id=skill.id, permission=permission, user_group_ids=user_group_ids, db=db, - ) - ] + ): + result.append(skill) + return result - def search_skills( + async def search_skills( self, user_id: str, filter: dict = {}, skip: int = 0, limit: int = 30, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> SkillListResponse: try: - with get_db_context(db) as db: - from open_webui.models.users import User, UserModel - + async with get_async_db_context(db) as db: # Join with User table for user filtering - query = db.query(Skill, User).outerjoin(User, User.id == Skill.user_id) + stmt = select(Skill, User).outerjoin(User, User.id == Skill.user_id) if filter: query_key = filter.get('query') if query_key: - query = query.filter( + stmt = stmt.filter( or_( Skill.name.ilike(f'%{query_key}%'), Skill.description.ilike(f'%{query_key}%'), @@ -242,44 +245,48 @@ class SkillsTable: view_option = filter.get('view_option') if view_option == 'created': - query = query.filter(Skill.user_id == user_id) + stmt = stmt.filter(Skill.user_id == user_id) elif view_option == 'shared': - query = query.filter(Skill.user_id != user_id) + stmt = stmt.filter(Skill.user_id != user_id) # Apply access grant filtering - query = AccessGrants.has_permission_filter( + stmt = AccessGrants.has_permission_filter( db=db, - query=query, + query=stmt, DocumentModel=Skill, filter=filter, resource_type='skill', permission='read', ) - query = query.order_by(Skill.updated_at.desc()) + stmt = stmt.order_by(Skill.updated_at.desc()) # Count BEFORE pagination - total = query.count() + count_result = await db.execute( + select(func.count()).select_from(stmt.subquery()) + ) + total = count_result.scalar() if skip: - query = query.offset(skip) + stmt = stmt.offset(skip) if limit: - query = query.limit(limit) + stmt = stmt.limit(limit) - items = query.all() + result = await db.execute(stmt) + items = result.all() skill_ids = [skill.id for skill, _ in items] - grants_map = AccessGrants.get_grants_by_resources('skill', skill_ids, db=db) + grants_map = await AccessGrants.get_grants_by_resources('skill', skill_ids, db=db) skills = [] for skill, user in items: skills.append( SkillUserResponse( - **self._to_skill_model( + **(await self._to_skill_model( skill, access_grants=grants_map.get(skill.id, []), db=db, - ).model_dump(), + )).model_dump(), user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None), ) ) @@ -289,43 +296,44 @@ class SkillsTable: log.exception(f'Error searching skills: {e}') return SkillListResponse(items=[], total=0) - def update_skill_by_id(self, id: str, updated: dict, db: Optional[Session] = None) -> Optional[SkillModel]: + async def update_skill_by_id(self, id: str, updated: dict, db: Optional[AsyncSession] = None) -> Optional[SkillModel]: try: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: access_grants = updated.pop('access_grants', None) - db.query(Skill).filter_by(id=id).update({**updated, 'updated_at': int(time.time())}) - db.commit() + await db.execute(update(Skill).filter_by(id=id).values(**updated, updated_at=int(time.time()))) + await db.commit() if access_grants is not None: - AccessGrants.set_access_grants('skill', id, access_grants, db=db) + await AccessGrants.set_access_grants('skill', id, access_grants, db=db) - skill = db.query(Skill).get(id) - db.refresh(skill) - return self._to_skill_model(skill, db=db) + skill = await db.get(Skill, id) + await db.refresh(skill) + return await self._to_skill_model(skill, db=db) except Exception: return None - def toggle_skill_by_id(self, id: str, db: Optional[Session] = None) -> Optional[SkillModel]: - with get_db_context(db) as db: + async def toggle_skill_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[SkillModel]: + async with get_async_db_context(db) as db: try: - skill = db.query(Skill).filter_by(id=id).first() + result = await db.execute(select(Skill).filter_by(id=id)) + skill = result.scalars().first() if not skill: return None skill.is_active = not skill.is_active skill.updated_at = int(time.time()) - db.commit() - db.refresh(skill) + await db.commit() + await db.refresh(skill) - return self._to_skill_model(skill, db=db) + return await self._to_skill_model(skill, db=db) except Exception: return None - def delete_skill_by_id(self, id: str, db: Optional[Session] = None) -> bool: + async def delete_skill_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: - AccessGrants.revoke_all_access('skill', id, db=db) - db.query(Skill).filter_by(id=id).delete() - db.commit() + async with get_async_db_context(db) as db: + await AccessGrants.revoke_all_access('skill', id, db=db) + await db.execute(delete(Skill).filter_by(id=id)) + await db.commit() return True except Exception: diff --git a/backend/open_webui/models/tags.py b/backend/open_webui/models/tags.py index b60220bc23..95b97b9cc1 100644 --- a/backend/open_webui/models/tags.py +++ b/backend/open_webui/models/tags.py @@ -3,8 +3,9 @@ import time import uuid from typing import Optional -from sqlalchemy.orm import Session -from open_webui.internal.db import Base, JSONField, get_db, get_db_context +from sqlalchemy import select, delete +from sqlalchemy.ext.asyncio import AsyncSession +from open_webui.internal.db import Base, JSONField, get_async_db_context from pydantic import BaseModel, ConfigDict @@ -53,15 +54,15 @@ class TagChatIdForm(BaseModel): class TagTable: - def insert_new_tag(self, name: str, user_id: str, db: Optional[Session] = None) -> Optional[TagModel]: - with get_db_context(db) as db: + async def insert_new_tag(self, name: str, user_id: str, db: Optional[AsyncSession] = None) -> Optional[TagModel]: + async with get_async_db_context(db) as db: id = name.replace(' ', '_').lower() tag = TagModel(**{'id': id, 'user_id': user_id, 'name': name}) try: result = Tag(**tag.model_dump()) db.add(result) - db.commit() - db.refresh(result) + await db.commit() + await db.refresh(result) if result: return TagModel.model_validate(result) else: @@ -70,64 +71,65 @@ class TagTable: log.exception(f'Error inserting a new tag: {e}') return None - def get_tag_by_name_and_user_id(self, name: str, user_id: str, db: Optional[Session] = None) -> Optional[TagModel]: + async def get_tag_by_name_and_user_id(self, name: str, user_id: str, db: Optional[AsyncSession] = None) -> Optional[TagModel]: try: id = name.replace(' ', '_').lower() - with get_db_context(db) as db: - tag = db.query(Tag).filter_by(id=id, user_id=user_id).first() - return TagModel.model_validate(tag) + async with get_async_db_context(db) as db: + result = await db.execute(select(Tag).filter_by(id=id, user_id=user_id)) + tag = result.scalars().first() + return TagModel.model_validate(tag) if tag else None except Exception: return None - def get_tags_by_user_id(self, user_id: str, db: Optional[Session] = None) -> list[TagModel]: - with get_db_context(db) as db: - return [TagModel.model_validate(tag) for tag in (db.query(Tag).filter_by(user_id=user_id).all())] + async def get_tags_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[TagModel]: + async with get_async_db_context(db) as db: + result = await db.execute(select(Tag).filter_by(user_id=user_id)) + return [TagModel.model_validate(tag) for tag in result.scalars().all()] - def get_tags_by_ids_and_user_id(self, ids: list[str], user_id: str, db: Optional[Session] = None) -> list[TagModel]: - with get_db_context(db) as db: - return [ - TagModel.model_validate(tag) - for tag in (db.query(Tag).filter(Tag.id.in_(ids), Tag.user_id == user_id).all()) - ] + async def get_tags_by_ids_and_user_id(self, ids: list[str], user_id: str, db: Optional[AsyncSession] = None) -> list[TagModel]: + async with get_async_db_context(db) as db: + result = await db.execute(select(Tag).filter(Tag.id.in_(ids), Tag.user_id == user_id)) + return [TagModel.model_validate(tag) for tag in result.scalars().all()] - def delete_tag_by_name_and_user_id(self, name: str, user_id: str, db: Optional[Session] = None) -> bool: + async def delete_tag_by_name_and_user_id(self, name: str, user_id: str, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: id = name.replace(' ', '_').lower() - res = db.query(Tag).filter_by(id=id, user_id=user_id).delete() - log.debug(f'res: {res}') - db.commit() + result = await db.execute(delete(Tag).filter_by(id=id, user_id=user_id)) + log.debug(f'res: {result.rowcount}') + await db.commit() return True except Exception as e: log.error(f'delete_tag: {e}') return False - def delete_tags_by_ids_and_user_id(self, ids: list[str], user_id: str, db: Optional[Session] = None) -> bool: + async def delete_tags_by_ids_and_user_id(self, ids: list[str], user_id: str, db: Optional[AsyncSession] = None) -> bool: """Delete all tags whose id is in *ids* for the given user, in one query.""" if not ids: return True try: - with get_db_context(db) as db: - db.query(Tag).filter(Tag.id.in_(ids), Tag.user_id == user_id).delete(synchronize_session=False) - db.commit() + async with get_async_db_context(db) as db: + await db.execute(delete(Tag).filter(Tag.id.in_(ids), Tag.user_id == user_id)) + await db.commit() return True except Exception as e: log.error(f'delete_tags_by_ids: {e}') return False - def ensure_tags_exist(self, names: list[str], user_id: str, db: Optional[Session] = None) -> None: + async def ensure_tags_exist(self, names: list[str], user_id: str, db: Optional[AsyncSession] = None) -> None: """Create tag rows for any *names* that don't already exist for *user_id*.""" if not names: return ids = [n.replace(' ', '_').lower() for n in names] - with get_db_context(db) as db: - existing = {t.id for t in db.query(Tag.id).filter(Tag.id.in_(ids), Tag.user_id == user_id).all()} + async with get_async_db_context(db) as db: + result = await db.execute(select(Tag.id).filter(Tag.id.in_(ids), Tag.user_id == user_id)) + existing = {row[0] for row in result.all()} new_tags = [ Tag(id=tag_id, name=name, user_id=user_id) for tag_id, name in zip(ids, names) if tag_id not in existing ] if new_tags: db.add_all(new_tags) - db.commit() + await db.commit() Tags = TagTable() diff --git a/backend/open_webui/models/tools.py b/backend/open_webui/models/tools.py index f89b98c5e7..fe772c4443 100644 --- a/backend/open_webui/models/tools.py +++ b/backend/open_webui/models/tools.py @@ -2,8 +2,9 @@ import logging import time from typing import Optional -from sqlalchemy.orm import Session, defer -from open_webui.internal.db import Base, JSONField, get_db, get_db_context +from sqlalchemy import select, delete, update +from sqlalchemy.ext.asyncio import AsyncSession +from open_webui.internal.db import Base, JSONField, get_async_db_context from open_webui.models.users import Users, UserResponse from open_webui.models.groups import Groups from open_webui.models.access_grants import AccessGrantModel, AccessGrants @@ -97,29 +98,29 @@ class ToolValves(BaseModel): class ToolsTable: - def _get_access_grants(self, tool_id: str, db: Optional[Session] = None) -> list[AccessGrantModel]: - return AccessGrants.get_grants_by_resource('tool', tool_id, db=db) + async def _get_access_grants(self, tool_id: str, db: Optional[AsyncSession] = None) -> list[AccessGrantModel]: + return await AccessGrants.get_grants_by_resource('tool', tool_id, db=db) - def _to_tool_model( + async def _to_tool_model( self, tool: Tool, access_grants: Optional[list[AccessGrantModel]] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> ToolModel: tool_data = ToolModel.model_validate(tool).model_dump(exclude={'access_grants'}) tool_data['access_grants'] = ( - access_grants if access_grants is not None else self._get_access_grants(tool_data['id'], db=db) + access_grants if access_grants is not None else await self._get_access_grants(tool_data['id'], db=db) ) return ToolModel.model_validate(tool_data) - def insert_new_tool( + async def insert_new_tool( self, user_id: str, form_data: ToolForm, specs: list[dict], - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[ToolModel]: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: try: result = Tool( **{ @@ -131,38 +132,39 @@ class ToolsTable: } ) db.add(result) - db.commit() - db.refresh(result) - AccessGrants.set_access_grants('tool', result.id, form_data.access_grants, db=db) + await db.commit() + await db.refresh(result) + await AccessGrants.set_access_grants('tool', result.id, form_data.access_grants, db=db) if result: - return self._to_tool_model(result, db=db) + return await self._to_tool_model(result, db=db) else: return None except Exception as e: log.exception(f'Error creating a new tool: {e}') return None - def get_tool_by_id(self, id: str, db: Optional[Session] = None) -> Optional[ToolModel]: + async def get_tool_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[ToolModel]: try: - with get_db_context(db) as db: - tool = db.get(Tool, id) - return self._to_tool_model(tool, db=db) if tool else None + async with get_async_db_context(db) as db: + tool = await db.get(Tool, id) + return await self._to_tool_model(tool, db=db) if tool else None except Exception: return None - def get_tools(self, defer_content: bool = False, db: Optional[Session] = None) -> list[ToolUserModel]: - with get_db_context(db) as db: - query = db.query(Tool).order_by(Tool.updated_at.desc()) + async def get_tools(self, defer_content: bool = False, db: Optional[AsyncSession] = None) -> list[ToolUserModel]: + async with get_async_db_context(db) as db: + stmt = select(Tool).order_by(Tool.updated_at.desc()) if defer_content: - query = query.options(defer(Tool.content), defer(Tool.specs)) - all_tools = query.all() + stmt = stmt + result = await db.execute(stmt) + all_tools = result.scalars().all() user_ids = list(set(tool.user_id for tool in all_tools)) tool_ids = [tool.id for tool in all_tools] - users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else [] + users = await Users.get_users_by_user_ids(user_ids, db=db) if user_ids else [] users_dict = {user.id: user for user in users} - grants_map = AccessGrants.get_grants_by_resources('tool', tool_ids, db=db) + grants_map = await AccessGrants.get_grants_by_resources('tool', tool_ids, db=db) tools = [] for tool in all_tools: @@ -170,62 +172,66 @@ class ToolsTable: tools.append( ToolUserModel.model_validate( { - **self._to_tool_model( + **(await self._to_tool_model( tool, access_grants=grants_map.get(tool.id, []), db=db, - ).model_dump(), + )).model_dump(), 'user': user.model_dump() if user else None, } ) ) return tools - def get_tools_by_user_id( + async def get_tools_by_user_id( self, user_id: str, permission: str = 'write', defer_content: bool = False, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> list[ToolUserModel]: - tools = self.get_tools(defer_content=defer_content, db=db) - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id, db=db)} + tools = await self.get_tools(defer_content=defer_content, db=db) + user_groups = await Groups.get_groups_by_member_id(user_id, db=db) + user_group_ids = {group.id for group in user_groups} - return [ - tool - for tool in tools - if tool.user_id == user_id - or AccessGrants.has_access( + result = [] + for tool in tools: + if tool.user_id == user_id: + result.append(tool) + elif await AccessGrants.has_access( user_id=user_id, resource_type='tool', resource_id=tool.id, permission=permission, user_group_ids=user_group_ids, db=db, - ) - ] + ): + result.append(tool) + return result - def get_tool_valves_by_id(self, id: str, db: Optional[Session] = None) -> Optional[dict]: + async def get_tool_valves_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[dict]: try: - with get_db_context(db) as db: - tool = db.get(Tool, id) + async with get_async_db_context(db) as db: + tool = await db.get(Tool, id) return tool.valves if tool.valves else {} except Exception as e: log.exception(f'Error getting tool valves by id {id}') return None - def update_tool_valves_by_id(self, id: str, valves: dict, db: Optional[Session] = None) -> Optional[ToolValves]: + async def update_tool_valves_by_id(self, id: str, valves: dict, db: Optional[AsyncSession] = None) -> Optional[ToolValves]: try: - with get_db_context(db) as db: - db.query(Tool).filter_by(id=id).update({'valves': valves, 'updated_at': int(time.time())}) - db.commit() - return self.get_tool_by_id(id, db=db) + async with get_async_db_context(db) as db: + await db.execute( + update(Tool).filter_by(id=id).values(valves=valves, updated_at=int(time.time())) + ) + await db.commit() + return await self.get_tool_by_id(id, db=db) except Exception: return None - def get_user_valves_by_id_and_user_id(self, id: str, user_id: str, db: Optional[Session] = None) -> Optional[dict]: + async def get_user_valves_by_id_and_user_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> Optional[dict]: try: - user = Users.get_user_by_id(user_id, db=db) + user = await Users.get_user_by_id(user_id, db=db) user_settings = user.settings.model_dump() if user.settings else {} # Check if user has "tools" and "valves" settings @@ -239,11 +245,11 @@ class ToolsTable: log.exception(f'Error getting user values by id {id} and user_id {user_id}: {e}') return None - def update_user_valves_by_id_and_user_id( - self, id: str, user_id: str, valves: dict, db: Optional[Session] = None + async def update_user_valves_by_id_and_user_id( + self, id: str, user_id: str, valves: dict, db: Optional[AsyncSession] = None ) -> Optional[dict]: try: - user = Users.get_user_by_id(user_id, db=db) + user = await Users.get_user_by_id(user_id, db=db) user_settings = user.settings.model_dump() if user.settings else {} # Check if user has "tools" and "valves" settings @@ -255,34 +261,36 @@ class ToolsTable: user_settings['tools']['valves'][id] = valves # Update the user settings in the database - Users.update_user_by_id(user_id, {'settings': user_settings}, db=db) + await Users.update_user_by_id(user_id, {'settings': user_settings}, db=db) return user_settings['tools']['valves'][id] except Exception as e: log.exception(f'Error updating user valves by id {id} and user_id {user_id}: {e}') return None - def update_tool_by_id(self, id: str, updated: dict, db: Optional[Session] = None) -> Optional[ToolModel]: + async def update_tool_by_id(self, id: str, updated: dict, db: Optional[AsyncSession] = None) -> Optional[ToolModel]: try: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: access_grants = updated.pop('access_grants', None) - db.query(Tool).filter_by(id=id).update({**updated, 'updated_at': int(time.time())}) - db.commit() + await db.execute( + update(Tool).filter_by(id=id).values(**updated, updated_at=int(time.time())) + ) + await db.commit() if access_grants is not None: - AccessGrants.set_access_grants('tool', id, access_grants, db=db) + await AccessGrants.set_access_grants('tool', id, access_grants, db=db) - tool = db.query(Tool).get(id) - db.refresh(tool) - return self._to_tool_model(tool, db=db) + tool = await db.get(Tool, id) + await db.refresh(tool) + return await self._to_tool_model(tool, db=db) except Exception: return None - def delete_tool_by_id(self, id: str, db: Optional[Session] = None) -> bool: + async def delete_tool_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: - AccessGrants.revoke_all_access('tool', id, db=db) - db.query(Tool).filter_by(id=id).delete() - db.commit() + async with get_async_db_context(db) as db: + await AccessGrants.revoke_all_access('tool', id, db=db) + await db.execute(delete(Tool).filter_by(id=id)) + await db.commit() return True except Exception: diff --git a/backend/open_webui/models/users.py b/backend/open_webui/models/users.py index ef90745efe..a5a43c27b8 100644 --- a/backend/open_webui/models/users.py +++ b/backend/open_webui/models/users.py @@ -1,20 +1,15 @@ import time from typing import Optional -from sqlalchemy.orm import Session, defer -from open_webui.internal.db import Base, JSONField, get_db, get_db_context - +from sqlalchemy import select, delete, update, func, or_, case, exists +from sqlalchemy.ext.asyncio import AsyncSession +from open_webui.internal.db import Base, JSONField, get_async_db_context from open_webui.env import DATABASE_USER_ACTIVE_STATUS_UPDATE_INTERVAL -from open_webui.models.chats import Chats -from open_webui.models.groups import Groups, GroupMember -from open_webui.models.channels import ChannelMember - from open_webui.utils.misc import throttle from open_webui.utils.validate import validate_profile_image_url - from pydantic import BaseModel, ConfigDict, field_validator, model_validator from sqlalchemy import ( BigInteger, @@ -24,11 +19,8 @@ from sqlalchemy import ( Boolean, Text, Date, - exists, - select, cast, ) -from sqlalchemy import or_, case, func from sqlalchemy.dialects.postgresql import JSONB import datetime @@ -39,13 +31,11 @@ import datetime # daily bread of every session. Let none go hungry. #################### - class UserSettings(BaseModel): ui: Optional[dict] = {} model_config = ConfigDict(extra='allow') pass - class User(Base): __tablename__ = 'user' @@ -79,7 +69,6 @@ class User(Base): updated_at = Column(BigInteger) created_at = Column(BigInteger) - class UserModel(BaseModel): id: str @@ -120,13 +109,11 @@ class UserModel(BaseModel): self.profile_image_url = f'/api/v1/users/{self.id}/profile/image' return self - class UserStatusModel(UserModel): is_active: bool = False model_config = ConfigDict(from_attributes=True) - class ApiKey(Base): __tablename__ = 'api_key' @@ -139,7 +126,6 @@ class ApiKey(Base): created_at = Column(BigInteger, nullable=False) updated_at = Column(BigInteger, nullable=False) - class ApiKeyModel(BaseModel): id: str user_id: str @@ -152,12 +138,10 @@ class ApiKeyModel(BaseModel): model_config = ConfigDict(from_attributes=True) - #################### # Forms #################### - class UpdateProfileForm(BaseModel): profile_image_url: str name: str @@ -170,31 +154,25 @@ class UpdateProfileForm(BaseModel): def check_profile_image_url(cls, v: str) -> str: return validate_profile_image_url(v) - class UserGroupIdsModel(UserModel): group_ids: list[str] = [] - class UserModelResponse(UserModel): model_config = ConfigDict(extra='allow') - class UserListResponse(BaseModel): users: list[UserModelResponse] total: int - class UserGroupIdsListResponse(BaseModel): users: list[UserGroupIdsModel] total: int - class UserStatus(BaseModel): status_emoji: Optional[str] = None status_message: Optional[str] = None status_expires_at: Optional[int] = None - class UserInfoResponse(UserStatus): id: str name: str @@ -204,48 +182,39 @@ class UserInfoResponse(UserStatus): groups: Optional[list] = [] is_active: bool = False - class UserIdNameResponse(BaseModel): id: str name: str - class UserIdNameStatusResponse(UserStatus): id: str name: str is_active: Optional[bool] = None - class UserInfoListResponse(BaseModel): users: list[UserInfoResponse] total: int - class UserIdNameListResponse(BaseModel): users: list[UserIdNameResponse] total: int - class UserNameResponse(BaseModel): id: str name: str role: str - class UserResponse(UserNameResponse): email: str - class UserProfileImageResponse(UserNameResponse): email: str profile_image_url: str - class UserRoleUpdateForm(BaseModel): id: str role: str - class UserUpdateForm(BaseModel): role: str name: str @@ -258,9 +227,8 @@ class UserUpdateForm(BaseModel): def check_profile_image_url(cls, v: str) -> str: return validate_profile_image_url(v) - class UsersTable: - def insert_new_user( + async def insert_new_user( self, id: str, name: str, @@ -269,9 +237,9 @@ class UsersTable: role: str = 'pending', username: Optional[str] = None, oauth: Optional[dict] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[UserModel]: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: user = UserModel( **{ 'id': id, @@ -288,87 +256,98 @@ class UsersTable: ) result = User(**user.model_dump()) db.add(result) - db.commit() - db.refresh(result) + await db.commit() + await db.refresh(result) if result: return user else: return None - def get_user_by_id(self, id: str, db: Optional[Session] = None) -> Optional[UserModel]: + async def get_user_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[UserModel]: try: - with get_db_context(db) as db: - user = db.query(User).filter_by(id=id).first() - return UserModel.model_validate(user) - except Exception: - return None - - def get_user_by_api_key(self, api_key: str, db: Optional[Session] = None) -> Optional[UserModel]: - try: - with get_db_context(db) as db: - user = db.query(User).join(ApiKey, User.id == ApiKey.user_id).filter(ApiKey.key == api_key).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(User).filter_by(id=id)) + user = result.scalars().first() return UserModel.model_validate(user) if user else None except Exception: return None - def get_user_by_email(self, email: str, db: Optional[Session] = None) -> Optional[UserModel]: + async def get_user_by_api_key(self, api_key: str, db: Optional[AsyncSession] = None) -> Optional[UserModel]: try: - with get_db_context(db) as db: - user = db.query(User).filter(func.lower(User.email) == email.lower()).first() + async with get_async_db_context(db) as db: + result = await db.execute( + select(User).join(ApiKey, User.id == ApiKey.user_id).filter(ApiKey.key == api_key) + ) + user = result.scalars().first() return UserModel.model_validate(user) if user else None except Exception: return None - def get_user_by_oauth_sub(self, provider: str, sub: str, db: Optional[Session] = None) -> Optional[UserModel]: + async def get_user_by_email(self, email: str, db: Optional[AsyncSession] = None) -> Optional[UserModel]: try: - with get_db_context(db) as db: # type: Session + async with get_async_db_context(db) as db: + result = await db.execute(select(User).filter(func.lower(User.email) == email.lower())) + user = result.scalars().first() + return UserModel.model_validate(user) if user else None + except Exception: + return None + + async def get_user_by_oauth_sub(self, provider: str, sub: str, db: Optional[AsyncSession] = None) -> Optional[UserModel]: + try: + async with get_async_db_context(db) as db: dialect_name = db.bind.dialect.name - query = db.query(User) + stmt = select(User) if dialect_name == 'sqlite': - query = query.filter(User.oauth.contains({provider: {'sub': sub}})) + stmt = stmt.filter(User.oauth.contains({provider: {'sub': sub}})) elif dialect_name == 'postgresql': - query = query.filter(User.oauth[provider].cast(JSONB)['sub'].astext == sub) + stmt = stmt.filter(User.oauth[provider].cast(JSONB)['sub'].astext == sub) - user = query.first() + result = await db.execute(stmt) + user = result.scalars().first() return UserModel.model_validate(user) if user else None except Exception as e: # You may want to log the exception here return None - def get_user_by_scim_external_id( - self, provider: str, external_id: str, db: Optional[Session] = None + async def get_user_by_scim_external_id( + self, provider: str, external_id: str, db: Optional[AsyncSession] = None ) -> Optional[UserModel]: try: - with get_db_context(db) as db: # type: Session + async with get_async_db_context(db) as db: dialect_name = db.bind.dialect.name - query = db.query(User) + stmt = select(User) if dialect_name == 'sqlite': - query = query.filter(User.scim.contains({provider: {'external_id': external_id}})) + stmt = stmt.filter(User.scim.contains({provider: {'external_id': external_id}})) elif dialect_name == 'postgresql': - query = query.filter(User.scim[provider].cast(JSONB)['external_id'].astext == external_id) + stmt = stmt.filter(User.scim[provider].cast(JSONB)['external_id'].astext == external_id) - user = query.first() + result = await db.execute(stmt) + user = result.scalars().first() return UserModel.model_validate(user) if user else None except Exception: return None - def get_users( + async def get_users( self, filter: Optional[dict] = None, skip: Optional[int] = None, limit: Optional[int] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> dict: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: + # Import here to avoid circular imports + from open_webui.models.groups import GroupMember + from open_webui.models.channels import ChannelMember + # Join GroupMember so we can order by group_id when requested - query = db.query(User).options(defer(User.profile_image_url)) + stmt = select(User) if filter: query_key = filter.get('query') if query_key: - query = query.filter( + stmt = stmt.filter( or_( User.name.ilike(f'%{query_key}%'), User.email.ilike(f'%{query_key}%'), @@ -377,7 +356,7 @@ class UsersTable: channel_id = filter.get('channel_id') if channel_id: - query = query.filter( + stmt = stmt.filter( exists( select(ChannelMember.id).where( ChannelMember.user_id == User.id, @@ -395,10 +374,10 @@ class UsersTable: return {'users': [], 'total': 0} if user_ids: - query = query.filter(User.id.in_(user_ids)) + stmt = stmt.filter(User.id.in_(user_ids)) if group_ids: - query = query.filter( + stmt = stmt.filter( exists( select(GroupMember.id).where( GroupMember.user_id == User.id, @@ -413,9 +392,9 @@ class UsersTable: exclude_roles = [role[1:] for role in roles if role.startswith('!')] if include_roles: - query = query.filter(User.role.in_(include_roles)) + stmt = stmt.filter(User.role.in_(include_roles)) if exclude_roles: - query = query.filter(~User.role.in_(exclude_roles)) + stmt = stmt.filter(~User.role.in_(exclude_roles)) order_by = filter.get('order_by') direction = filter.get('direction') @@ -435,99 +414,111 @@ class UsersTable: group_sort = case((membership_exists, 1), else_=0) if direction == 'asc': - query = query.order_by(group_sort.asc(), User.name.asc()) + stmt = stmt.order_by(group_sort.asc(), User.name.asc()) else: - query = query.order_by(group_sort.desc(), User.name.asc()) + stmt = stmt.order_by(group_sort.desc(), User.name.asc()) elif order_by == 'name': if direction == 'asc': - query = query.order_by(User.name.asc()) + stmt = stmt.order_by(User.name.asc()) else: - query = query.order_by(User.name.desc()) + stmt = stmt.order_by(User.name.desc()) elif order_by == 'email': if direction == 'asc': - query = query.order_by(User.email.asc()) + stmt = stmt.order_by(User.email.asc()) else: - query = query.order_by(User.email.desc()) + stmt = stmt.order_by(User.email.desc()) elif order_by == 'created_at': if direction == 'asc': - query = query.order_by(User.created_at.asc()) + stmt = stmt.order_by(User.created_at.asc()) else: - query = query.order_by(User.created_at.desc()) + stmt = stmt.order_by(User.created_at.desc()) elif order_by == 'last_active_at': if direction == 'asc': - query = query.order_by(User.last_active_at.asc()) + stmt = stmt.order_by(User.last_active_at.asc()) else: - query = query.order_by(User.last_active_at.desc()) + stmt = stmt.order_by(User.last_active_at.desc()) elif order_by == 'updated_at': if direction == 'asc': - query = query.order_by(User.updated_at.asc()) + stmt = stmt.order_by(User.updated_at.asc()) else: - query = query.order_by(User.updated_at.desc()) + stmt = stmt.order_by(User.updated_at.desc()) elif order_by == 'role': if direction == 'asc': - query = query.order_by(User.role.asc()) + stmt = stmt.order_by(User.role.asc()) else: - query = query.order_by(User.role.desc()) + stmt = stmt.order_by(User.role.desc()) else: - query = query.order_by(User.created_at.desc()) + stmt = stmt.order_by(User.created_at.desc()) # Count BEFORE pagination - total = query.count() + count_result = await db.execute( + select(func.count()).select_from(stmt.subquery()) + ) + total = count_result.scalar() # correct pagination logic if skip is not None: - query = query.offset(skip) + stmt = stmt.offset(skip) if limit is not None: - query = query.limit(limit) + stmt = stmt.limit(limit) - users = query.all() + result = await db.execute(stmt) + users = result.scalars().all() return { 'users': [UserModel.model_validate(user) for user in users], 'total': total, } - def get_users_by_group_id(self, group_id: str, db: Optional[Session] = None) -> list[UserModel]: - with get_db_context(db) as db: - users = ( - db.query(User) - .options(defer(User.profile_image_url)) + async def get_users_by_group_id(self, group_id: str, db: Optional[AsyncSession] = None) -> list[UserModel]: + async with get_async_db_context(db) as db: + from open_webui.models.groups import GroupMember + result = await db.execute( + select(User) + .join(GroupMember, User.id == GroupMember.user_id) .filter(GroupMember.group_id == group_id) - .all() ) + users = result.scalars().all() return [UserModel.model_validate(user) for user in users] - def get_users_by_user_ids(self, user_ids: list[str], db: Optional[Session] = None) -> list[UserStatusModel]: - with get_db_context(db) as db: - users = db.query(User).options(defer(User.profile_image_url)).filter(User.id.in_(user_ids)).all() + async def get_users_by_user_ids(self, user_ids: list[str], db: Optional[AsyncSession] = None) -> list[UserStatusModel]: + async with get_async_db_context(db) as db: + result = await db.execute( + select(User).filter(User.id.in_(user_ids)) + ) + users = result.scalars().all() return [UserModel.model_validate(user) for user in users] - def get_num_users(self, db: Optional[Session] = None) -> Optional[int]: - with get_db_context(db) as db: - return db.query(User).count() + async def get_num_users(self, db: Optional[AsyncSession] = None) -> Optional[int]: + async with get_async_db_context(db) as db: + result = await db.execute(select(func.count()).select_from(User)) + return result.scalar() - def has_users(self, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: - return db.query(db.query(User).exists()).scalar() + async def has_users(self, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: + result = await db.execute(select(exists(select(User)))) + return result.scalar() - def get_first_user(self, db: Optional[Session] = None) -> UserModel: + async def get_first_user(self, db: Optional[AsyncSession] = None) -> UserModel: try: - with get_db_context(db) as db: - user = db.query(User).order_by(User.created_at).first() - return UserModel.model_validate(user) + async with get_async_db_context(db) as db: + result = await db.execute(select(User).order_by(User.created_at).limit(1)) + user = result.scalars().first() + return UserModel.model_validate(user) if user else None except Exception: return None - def get_user_webhook_url_by_id(self, id: str, db: Optional[Session] = None) -> Optional[str]: + async def get_user_webhook_url_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[str]: try: - with get_db_context(db) as db: - user = db.query(User).filter_by(id=id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(User).filter_by(id=id)) + user = result.scalars().first() if user.settings is None: return None @@ -536,68 +527,73 @@ class UsersTable: except Exception: return None - def get_num_users_active_today(self, db: Optional[Session] = None) -> Optional[int]: - with get_db_context(db) as db: + async def get_num_users_active_today(self, db: Optional[AsyncSession] = None) -> Optional[int]: + async with get_async_db_context(db) as db: current_timestamp = int(datetime.datetime.now().timestamp()) today_midnight_timestamp = current_timestamp - (current_timestamp % 86400) - query = db.query(User).filter(User.last_active_at > today_midnight_timestamp) - return query.count() + result = await db.execute( + select(func.count()).select_from(User).filter(User.last_active_at > today_midnight_timestamp) + ) + return result.scalar() - def update_user_role_by_id(self, id: str, role: str, db: Optional[Session] = None) -> Optional[UserModel]: + async def update_user_role_by_id(self, id: str, role: str, db: Optional[AsyncSession] = None) -> Optional[UserModel]: try: - with get_db_context(db) as db: - user = db.query(User).filter_by(id=id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(User).filter_by(id=id)) + user = result.scalars().first() if not user: return None user.role = role - db.commit() - db.refresh(user) + await db.commit() + await db.refresh(user) return UserModel.model_validate(user) except Exception: return None - def update_user_status_by_id( - self, id: str, form_data: UserStatus, db: Optional[Session] = None + async def update_user_status_by_id( + self, id: str, form_data: UserStatus, db: Optional[AsyncSession] = None ) -> Optional[UserModel]: try: - with get_db_context(db) as db: - user = db.query(User).filter_by(id=id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(User).filter_by(id=id)) + user = result.scalars().first() if not user: return None for key, value in form_data.model_dump(exclude_none=True).items(): setattr(user, key, value) - db.commit() - db.refresh(user) + await db.commit() + await db.refresh(user) return UserModel.model_validate(user) except Exception: return None - def update_user_profile_image_url_by_id( - self, id: str, profile_image_url: str, db: Optional[Session] = None + async def update_user_profile_image_url_by_id( + self, id: str, profile_image_url: str, db: Optional[AsyncSession] = None ) -> Optional[UserModel]: try: - with get_db_context(db) as db: - user = db.query(User).filter_by(id=id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(User).filter_by(id=id)) + user = result.scalars().first() if not user: return None user.profile_image_url = profile_image_url - db.commit() - db.refresh(user) + await db.commit() + await db.refresh(user) return UserModel.model_validate(user) except Exception: return None @throttle(DATABASE_USER_ACTIVE_STATUS_UPDATE_INTERVAL) - def update_last_active_by_id(self, id: str, db: Optional[Session] = None) -> None: + async def update_last_active_by_id(self, id: str, db: Optional[AsyncSession] = None) -> None: try: - with get_db_context(db) as db: - db.query(User).filter_by(id=id).update({'last_active_at': int(time.time())}) - db.commit() + async with get_async_db_context(db) as db: + await db.execute(update(User).filter_by(id=id).values(last_active_at=int(time.time()))) + await db.commit() except Exception: pass - def update_user_oauth_by_id( - self, id: str, provider: str, sub: str, db: Optional[Session] = None + async def update_user_oauth_by_id( + self, id: str, provider: str, sub: str, db: Optional[AsyncSession] = None ) -> Optional[UserModel]: """ Update or insert an OAuth provider/sub pair into the user's oauth JSON field. @@ -608,8 +604,9 @@ class UsersTable: } """ try: - with get_db_context(db) as db: - user = db.query(User).filter_by(id=id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(User).filter_by(id=id)) + user = result.scalars().first() if not user: return None @@ -620,20 +617,20 @@ class UsersTable: oauth[provider] = {'sub': sub} # Persist updated JSON - db.query(User).filter_by(id=id).update({'oauth': oauth}) - db.commit() + await db.execute(update(User).filter_by(id=id).values(oauth=oauth)) + await db.commit() return UserModel.model_validate(user) except Exception: return None - def update_user_scim_by_id( + async def update_user_scim_by_id( self, id: str, provider: str, external_id: str, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> Optional[UserModel]: """ Update or insert a SCIM provider/external_id pair into the user's scim JSON field. @@ -644,41 +641,44 @@ class UsersTable: } """ try: - with get_db_context(db) as db: - user = db.query(User).filter_by(id=id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(User).filter_by(id=id)) + user = result.scalars().first() if not user: return None scim = user.scim or {} scim[provider] = {'external_id': external_id} - db.query(User).filter_by(id=id).update({'scim': scim}) - db.commit() + await db.execute(update(User).filter_by(id=id).values(scim=scim)) + await db.commit() return UserModel.model_validate(user) except Exception: return None - def update_user_by_id(self, id: str, updated: dict, db: Optional[Session] = None) -> Optional[UserModel]: + async def update_user_by_id(self, id: str, updated: dict, db: Optional[AsyncSession] = None) -> Optional[UserModel]: try: - with get_db_context(db) as db: - user = db.query(User).filter_by(id=id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(User).filter_by(id=id)) + user = result.scalars().first() if not user: return None for key, value in updated.items(): setattr(user, key, value) - db.commit() - db.refresh(user) + await db.commit() + await db.refresh(user) return UserModel.model_validate(user) except Exception as e: print(e) return None - def update_user_settings_by_id(self, id: str, updated: dict, db: Optional[Session] = None) -> Optional[UserModel]: + async def update_user_settings_by_id(self, id: str, updated: dict, db: Optional[AsyncSession] = None) -> Optional[UserModel]: try: - with get_db_context(db) as db: - user = db.query(User).filter_by(id=id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(User).filter_by(id=id)) + user = result.scalars().first() if not user: return None @@ -689,26 +689,30 @@ class UsersTable: user_settings.update(updated) - db.query(User).filter_by(id=id).update({'settings': user_settings}) - db.commit() + await db.execute(update(User).filter_by(id=id).values(settings=user_settings)) + await db.commit() - user = db.query(User).filter_by(id=id).first() + result = await db.execute(select(User).filter_by(id=id)) + user = result.scalars().first() return UserModel.model_validate(user) except Exception: return None - def delete_user_by_id(self, id: str, db: Optional[Session] = None) -> bool: + async def delete_user_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool: try: + from open_webui.models.groups import Groups + from open_webui.models.chats import Chats + # Remove User from Groups - Groups.remove_user_from_all_groups(id) + await Groups.remove_user_from_all_groups(id) # Delete User Chats - result = Chats.delete_chats_by_user_id(id, db=db) + result = await Chats.delete_chats_by_user_id(id, db=db) if result: - with get_db_context(db) as db: + async with get_async_db_context(db) as db: # Delete User - db.query(User).filter_by(id=id).delete() - db.commit() + await db.execute(delete(User).filter_by(id=id)) + await db.commit() return True else: @@ -716,19 +720,20 @@ class UsersTable: except Exception: return False - def get_user_api_key_by_id(self, id: str, db: Optional[Session] = None) -> Optional[str]: + async def get_user_api_key_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[str]: try: - with get_db_context(db) as db: - api_key = db.query(ApiKey).filter_by(user_id=id).first() + async with get_async_db_context(db) as db: + result = await db.execute(select(ApiKey).filter_by(user_id=id)) + api_key = result.scalars().first() return api_key.key if api_key else None except Exception: return None - def update_user_api_key_by_id(self, id: str, api_key: str, db: Optional[Session] = None) -> bool: + async def update_user_api_key_by_id(self, id: str, api_key: str, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: - db.query(ApiKey).filter_by(user_id=id).delete() - db.commit() + async with get_async_db_context(db) as db: + await db.execute(delete(ApiKey).filter_by(user_id=id)) + await db.commit() now = int(time.time()) new_api_key = ApiKey( @@ -739,41 +744,45 @@ class UsersTable: updated_at=now, ) db.add(new_api_key) - db.commit() + await db.commit() return True except Exception: return False - def delete_user_api_key_by_id(self, id: str, db: Optional[Session] = None) -> bool: + async def delete_user_api_key_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool: try: - with get_db_context(db) as db: - db.query(ApiKey).filter_by(user_id=id).delete() - db.commit() + async with get_async_db_context(db) as db: + await db.execute(delete(ApiKey).filter_by(user_id=id)) + await db.commit() return True except Exception: return False - def get_valid_user_ids(self, user_ids: list[str], db: Optional[Session] = None) -> list[str]: - with get_db_context(db) as db: - users = db.query(User).filter(User.id.in_(user_ids)).all() + async def get_valid_user_ids(self, user_ids: list[str], db: Optional[AsyncSession] = None) -> list[str]: + async with get_async_db_context(db) as db: + result = await db.execute(select(User).filter(User.id.in_(user_ids))) + users = result.scalars().all() return [user.id for user in users] - def get_super_admin_user(self, db: Optional[Session] = None) -> Optional[UserModel]: - with get_db_context(db) as db: - user = db.query(User).filter_by(role='admin').first() + async def get_super_admin_user(self, db: Optional[AsyncSession] = None) -> Optional[UserModel]: + async with get_async_db_context(db) as db: + result = await db.execute(select(User).filter_by(role='admin').limit(1)) + user = result.scalars().first() if user: return UserModel.model_validate(user) else: return None - def get_active_user_count(self, db: Optional[Session] = None) -> int: - with get_db_context(db) as db: + async def get_active_user_count(self, db: Optional[AsyncSession] = None) -> int: + async with get_async_db_context(db) as db: # Consider user active if last_active_at within the last 3 minutes three_minutes_ago = int(time.time()) - 180 - count = db.query(User).filter(User.last_active_at >= three_minutes_ago).count() - return count + result = await db.execute( + select(func.count()).select_from(User).filter(User.last_active_at >= three_minutes_ago) + ) + return result.scalar() @staticmethod def is_active(user: UserModel) -> bool: @@ -783,14 +792,14 @@ class UsersTable: return user.last_active_at >= three_minutes_ago return False - def is_user_active(self, user_id: str, db: Optional[Session] = None) -> bool: - with get_db_context(db) as db: - user = db.query(User).filter_by(id=user_id).first() + async def is_user_active(self, user_id: str, db: Optional[AsyncSession] = None) -> bool: + async with get_async_db_context(db) as db: + result = await db.execute(select(User).filter_by(id=user_id)) + user = result.scalars().first() if user and user.last_active_at: # Consider user active if last_active_at within the last 3 minutes three_minutes_ago = int(time.time()) - 180 return user.last_active_at >= three_minutes_ago return False - Users = UsersTable() diff --git a/backend/open_webui/retrieval/utils.py b/backend/open_webui/retrieval/utils.py index 4ab8bdf7c0..f7d2775c52 100644 --- a/backend/open_webui/retrieval/utils.py +++ b/backend/open_webui/retrieval/utils.py @@ -978,12 +978,12 @@ async def get_sources_from_items( elif item.get('type') == 'note': # Note Attached - note = Notes.get_note_by_id(item.get('id')) + note = await Notes.get_note_by_id(item.get('id')) if note and ( user.role == 'admin' or note.user_id == user.id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='note', resource_id=note.id, @@ -998,7 +998,7 @@ async def get_sources_from_items( elif item.get('type') == 'chat': # Chat Attached - chat = Chats.get_chat_by_id(item.get('id')) + chat = await Chats.get_chat_by_id(item.get('id')) if chat and (user.role == 'admin' or chat.user_id == user.id): messages_map = chat.chat.get('history', {}).get('messages', {}) @@ -1042,11 +1042,11 @@ async def get_sources_from_items( ], } elif item.get('id'): - file_object = Files.get_file_by_id(item.get('id')) + file_object = await Files.get_file_by_id(item.get('id')) if file_object and ( user.role == 'admin' or file_object.user_id == user.id - or has_access_to_file(item.get('id'), 'read', user) + or await has_access_to_file(item.get('id'), 'read', user) ): query_result = { 'documents': [[file_object.data.get('content', '')]], @@ -1069,12 +1069,12 @@ async def get_sources_from_items( elif item.get('type') == 'collection': # Manual Full Mode Toggle for Collection - knowledge_base = Knowledges.get_knowledge_by_id(item.get('id')) + knowledge_base = await Knowledges.get_knowledge_by_id(item.get('id')) if knowledge_base and ( user.role == 'admin' or knowledge_base.user_id == user.id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='knowledge', resource_id=knowledge_base.id, @@ -1085,14 +1085,14 @@ async def get_sources_from_items( if knowledge_base and ( user.role == 'admin' or knowledge_base.user_id == user.id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='knowledge', resource_id=knowledge_base.id, permission='read', ) ): - files = Knowledges.get_files_by_id(knowledge_base.id) + files = await Knowledges.get_files_by_id(knowledge_base.id) documents = [] metadatas = [] diff --git a/backend/open_webui/routers/analytics.py b/backend/open_webui/routers/analytics.py index 790c134295..8636444a5d 100644 --- a/backend/open_webui/routers/analytics.py +++ b/backend/open_webui/routers/analytics.py @@ -11,8 +11,8 @@ from open_webui.models.groups import Groups from open_webui.models.users import Users from open_webui.models.feedbacks import Feedbacks from open_webui.utils.auth import get_admin_user -from open_webui.internal.db import get_session -from sqlalchemy.orm import Session +from open_webui.internal.db import get_async_session +from sqlalchemy.ext.asyncio import AsyncSession log = logging.getLogger(__name__) @@ -59,10 +59,10 @@ async def get_model_analytics( end_date: Optional[int] = Query(None, description='End timestamp (epoch)'), group_id: Optional[str] = Query(None, description='Filter by user group ID'), user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Get message counts per model.""" - counts = ChatMessages.get_message_count_by_model(start_date=start_date, end_date=end_date, group_id=group_id, db=db) + counts = await ChatMessages.get_message_count_by_model(start_date=start_date, end_date=end_date, group_id=group_id, db=db) models = [ ModelAnalyticsEntry(model_id=model_id, count=count) for model_id, count in sorted(counts.items(), key=lambda x: -x[1]) @@ -77,17 +77,17 @@ async def get_user_analytics( group_id: Optional[str] = Query(None, description='Filter by user group ID'), limit: int = Query(50, description='Max users to return'), user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Get message counts and token usage per user with user info.""" - counts = ChatMessages.get_message_count_by_user(start_date=start_date, end_date=end_date, group_id=group_id, db=db) - token_usage = ChatMessages.get_token_usage_by_user( + counts = await ChatMessages.get_message_count_by_user(start_date=start_date, end_date=end_date, group_id=group_id, db=db) + token_usage = await ChatMessages.get_token_usage_by_user( start_date=start_date, end_date=end_date, group_id=group_id, db=db ) # Get user info for top users top_user_ids = [uid for uid, _ in sorted(counts.items(), key=lambda x: -x[1])[:limit]] - user_info = {u.id: u for u in Users.get_users_by_user_ids(top_user_ids, db=db)} + user_info = {u.id: u for u in await Users.get_users_by_user_ids(top_user_ids, db=db)} users = [] for user_id in top_user_ids: @@ -118,13 +118,13 @@ async def get_messages( skip: int = Query(0), limit: int = Query(50, le=100), user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Query messages with filters.""" if chat_id: - return ChatMessages.get_messages_by_chat_id(chat_id=chat_id, db=db) + return await ChatMessages.get_messages_by_chat_id(chat_id=chat_id, db=db) elif model_id: - return ChatMessages.get_messages_by_model_id( + return await ChatMessages.get_messages_by_model_id( model_id=model_id, start_date=start_date, end_date=end_date, @@ -133,7 +133,7 @@ async def get_messages( db=db, ) elif user_id: - return ChatMessages.get_messages_by_user_id(user_id=user_id, skip=skip, limit=limit, db=db) + return await ChatMessages.get_messages_by_user_id(user_id=user_id, skip=skip, limit=limit, db=db) else: # Return empty if no filter specified return [] @@ -152,16 +152,16 @@ async def get_summary( end_date: Optional[int] = Query(None, description='End timestamp (epoch)'), group_id: Optional[str] = Query(None, description='Filter by user group ID'), user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Get summary statistics for the dashboard.""" - model_counts = ChatMessages.get_message_count_by_model( + model_counts = await ChatMessages.get_message_count_by_model( start_date=start_date, end_date=end_date, group_id=group_id, db=db ) - user_counts = ChatMessages.get_message_count_by_user( + user_counts = await ChatMessages.get_message_count_by_user( start_date=start_date, end_date=end_date, group_id=group_id, db=db ) - chat_counts = ChatMessages.get_message_count_by_chat( + chat_counts = await ChatMessages.get_message_count_by_chat( start_date=start_date, end_date=end_date, group_id=group_id, db=db ) @@ -189,13 +189,13 @@ async def get_daily_stats( group_id: Optional[str] = Query(None, description='Filter by user group ID'), granularity: str = Query('daily', description="Granularity: 'hourly' or 'daily'"), user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Get message counts grouped by model for time-series chart.""" if granularity == 'hourly': - counts = ChatMessages.get_hourly_message_counts_by_model(start_date=start_date, end_date=end_date, db=db) + counts = await ChatMessages.get_hourly_message_counts_by_model(start_date=start_date, end_date=end_date, db=db) else: - counts = ChatMessages.get_daily_message_counts_by_model( + counts = await ChatMessages.get_daily_message_counts_by_model( start_date=start_date, end_date=end_date, group_id=group_id, db=db ) return DailyStatsResponse( @@ -224,10 +224,10 @@ async def get_token_usage( end_date: Optional[int] = Query(None), group_id: Optional[str] = Query(None, description='Filter by user group ID'), user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Get token usage aggregated by model.""" - usage = ChatMessages.get_token_usage_by_model(start_date=start_date, end_date=end_date, group_id=group_id, db=db) + usage = await ChatMessages.get_token_usage_by_model(start_date=start_date, end_date=end_date, group_id=group_id, db=db) models = [ TokenUsageEntry(model_id=model_id, **data) @@ -271,12 +271,12 @@ async def get_model_chats( skip: int = Query(0), limit: int = Query(50, le=100), user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Get chats that used a specific model, with preview and feedback info.""" # Get chat IDs that used this model - chat_ids = ChatMessages.get_chat_ids_by_model_id( + chat_ids = await ChatMessages.get_chat_ids_by_model_id( model_id=model_id, start_date=start_date, end_date=end_date, @@ -291,7 +291,7 @@ async def get_model_chats( # Get chat details from messages only chats_data = [] for chat_id in chat_ids: - messages = ChatMessages.get_messages_by_chat_id(chat_id, db=db) + messages = await ChatMessages.get_messages_by_chat_id(chat_id, db=db) if not messages: continue @@ -312,7 +312,7 @@ async def get_model_chats( # Get user info user_name = None if user_id: - user_info = Users.get_user_by_id(user_id, db=db) + user_info = await Users.get_user_by_id(user_id, db=db) user_name = user_info.name if user_info else None # Timestamps from messages @@ -357,12 +357,12 @@ async def get_model_overview( model_id: str, days: int = Query(30, description='Number of days of history (0 for all)'), user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Get model overview with feedback history and chat tags.""" # Get chat IDs that used this model - chat_ids = ChatMessages.get_chat_ids_by_model_id( + chat_ids = await ChatMessages.get_chat_ids_by_model_id( model_id=model_id, start_date=None, end_date=None, @@ -381,7 +381,7 @@ async def get_model_overview( start_dt = now - timedelta(days=days) for chat_id in chat_ids: - feedbacks = Feedbacks.get_feedbacks_by_chat_id(chat_id, db=db) + feedbacks = await Feedbacks.get_feedbacks_by_chat_id(chat_id, db=db) for fb in feedbacks: if fb.data and 'rating' in fb.data: rating = fb.data['rating'] @@ -425,7 +425,7 @@ async def get_model_overview( # Get chat tags tag_counts: dict[str, int] = defaultdict(int) for chat_id in chat_ids: - chat = Chats.get_chat_by_id(chat_id, db=db) + chat = await Chats.get_chat_by_id(chat_id, db=db) if chat and chat.meta: for tag in chat.meta.get('tags', []): tag_counts[tag] += 1 diff --git a/backend/open_webui/routers/audio.py b/backend/open_webui/routers/audio.py index b1886b2137..5a26d04e0f 100644 --- a/backend/open_webui/routers/audio.py +++ b/backend/open_webui/routers/audio.py @@ -330,7 +330,7 @@ async def speech(request: Request, user=Depends(get_verified_user)): detail=ERROR_MESSAGES.NOT_FOUND, ) - if user.role != 'admin' and not has_permission(user.id, 'chat.tts', request.app.state.config.USER_PERMISSIONS): + if user.role != 'admin' and not await has_permission(user.id, 'chat.tts', request.app.state.config.USER_PERMISSIONS): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.ACCESS_PROHIBITED, @@ -1208,13 +1208,13 @@ def split_audio(file_path, max_bytes, format='mp3', bitrate='32k'): @router.post('/transcriptions') -def transcription( +async def transcription( request: Request, file: UploadFile = File(...), language: Optional[str] = Form(None), user=Depends(get_verified_user), ): - if user.role != 'admin' and not has_permission(user.id, 'chat.stt', request.app.state.config.USER_PERMISSIONS): + if user.role != 'admin' and not await has_permission(user.id, 'chat.stt', request.app.state.config.USER_PERMISSIONS): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.ACCESS_PROHIBITED, diff --git a/backend/open_webui/routers/auths.py b/backend/open_webui/routers/auths.py index 8d7581dd4b..484212a493 100644 --- a/backend/open_webui/routers/auths.py +++ b/backend/open_webui/routers/auths.py @@ -70,8 +70,8 @@ from open_webui.utils.auth import ( get_password_hash, get_http_authorization_cred, ) -from open_webui.internal.db import get_session -from sqlalchemy.orm import Session +from open_webui.internal.db import get_async_session +from sqlalchemy.ext.asyncio import AsyncSession from open_webui.utils.webhook import post_webhook from open_webui.utils.access_control import get_permissions, has_permission from open_webui.utils.groups import apply_default_group_assignment @@ -96,7 +96,7 @@ log = logging.getLogger(__name__) signin_rate_limiter = RateLimiter(redis_client=get_redis_client(), limit=5 * 3, window=60 * 3) -def create_session_response(request: Request, user, db, response: Response = None, set_cookie: bool = False) -> dict: +async def create_session_response(request: Request, user, db, response: Response = None, set_cookie: bool = False) -> dict: """ Create JWT token and build session response for a user. Shared helper for signin, signup, ldap_auth, add_user, and token_exchange endpoints. @@ -131,7 +131,7 @@ def create_session_response(request: Request, user, db, response: Response = Non **({'max_age': max_age} if max_age is not None else {}), ) - user_permissions = get_permissions(user.id, request.app.state.config.USER_PERMISSIONS, db=db) + user_permissions = await get_permissions(user.id, request.app.state.config.USER_PERMISSIONS, db=db) return { 'token': token, @@ -167,7 +167,7 @@ async def get_session_user( request: Request, response: Response, user=Depends(get_current_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): auth_header = request.headers.get('Authorization') auth_token = get_http_authorization_cred(auth_header) @@ -197,7 +197,7 @@ async def get_session_user( **({'max_age': max_age} if max_age is not None else {}), ) - user_permissions = get_permissions(user.id, request.app.state.config.USER_PERMISSIONS, db=db) + user_permissions = await get_permissions(user.id, request.app.state.config.USER_PERMISSIONS, db=db) return { 'token': token, @@ -227,10 +227,10 @@ async def get_session_user( async def update_profile( form_data: UpdateProfileForm, session_user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): if session_user: - user = Users.update_user_by_id( + user = await Users.update_user_by_id( session_user.id, form_data.model_dump(), db=db, @@ -256,10 +256,10 @@ class UpdateTimezoneForm(BaseModel): async def update_timezone( form_data: UpdateTimezoneForm, session_user=Depends(get_current_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): if session_user: - Users.update_user_by_id( + await Users.update_user_by_id( session_user.id, {'timezone': form_data.timezone}, db=db, @@ -278,12 +278,12 @@ async def update_timezone( async def update_password( form_data: UpdatePasswordForm, session_user=Depends(get_current_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): if WEBUI_AUTH_TRUSTED_EMAIL_HEADER: raise HTTPException(400, detail=ERROR_MESSAGES.ACTION_PROHIBITED) if session_user: - user = Auths.authenticate_user( + user = await Auths.authenticate_user( session_user.email, lambda pw: verify_password(form_data.password, pw), db=db, @@ -295,7 +295,7 @@ async def update_password( except Exception as e: raise HTTPException(400, detail=str(e)) hashed = get_password_hash(form_data.new_password) - return Auths.update_user_password_by_id(user.id, hashed, db=db) + return await Auths.update_user_password_by_id(user.id, hashed, db=db) else: raise HTTPException(400, detail=ERROR_MESSAGES.INCORRECT_PASSWORD) else: @@ -310,7 +310,7 @@ async def ldap_auth( request: Request, response: Response, form_data: LdapForm, - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): # Security checks FIRST - before loading any config if not request.app.state.config.ENABLE_LDAP: @@ -476,12 +476,12 @@ async def ldap_auth( if not await asyncio.to_thread(connection_user.bind): raise HTTPException(400, 'Authentication failed.') - user = Users.get_user_by_email(email, db=db) + user = await Users.get_user_by_email(email, db=db) if not user: try: # Insert with default role first to avoid TOCTOU race on # first-user registration. Matches signup_handler pattern. - user = Auths.insert_new_auth( + user = await Auths.insert_new_auth( email=email, password=str(uuid.uuid4()), name=cn, @@ -494,11 +494,11 @@ async def ldap_auth( # Atomically check if this is the only user *after* the # insert. Only the single user present should become admin. - if Users.get_num_users(db=db) == 1: - Users.update_user_role_by_id(user.id, 'admin', db=db) - user = Users.get_user_by_id(user.id, db=db) + if await Users.get_num_users(db=db) == 1: + await Users.update_user_role_by_id(user.id, 'admin', db=db) + user = await Users.get_user_by_id(user.id, db=db) - apply_default_group_assignment( + await apply_default_group_assignment( request.app.state.config.DEFAULT_GROUP_ID, user.id, db=db, @@ -510,19 +510,19 @@ async def ldap_auth( log.error(f'LDAP user creation error: {str(err)}') raise HTTPException(500, detail='Internal error occurred during LDAP user creation.') - user = Auths.authenticate_user_by_email(email, db=db) + user = await Auths.authenticate_user_by_email(email, db=db) if user: if ENABLE_LDAP_GROUP_MANAGEMENT and user_groups: if ENABLE_LDAP_GROUP_CREATION: - Groups.create_groups_by_group_names(user.id, user_groups, db=db) + await Groups.create_groups_by_group_names(user.id, user_groups, db=db) try: - Groups.sync_groups_by_group_names(user.id, user_groups, db=db) + await Groups.sync_groups_by_group_names(user.id, user_groups, db=db) log.info(f'Successfully synced groups for user {user.id}: {user_groups}') except Exception as e: log.error(f'Failed to sync groups for user {user.id}: {e}') - return create_session_response(request, user, db, response, set_cookie=True) + return await create_session_response(request, user, db, response, set_cookie=True) else: raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED) else: @@ -542,7 +542,7 @@ async def signin( request: Request, response: Response, form_data: SigninForm, - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): if not ENABLE_PASSWORD_AUTH: raise HTTPException( @@ -564,7 +564,7 @@ async def signin( except Exception as e: pass - if not Users.get_user_by_email(email.lower(), db=db): + if not await Users.get_user_by_email(email.lower(), db=db): await signup_handler( request, email, @@ -573,20 +573,20 @@ async def signin( db=db, ) - user = Auths.authenticate_user_by_email(email, db=db) + user = await Auths.authenticate_user_by_email(email, db=db) if user: if WEBUI_AUTH_TRUSTED_GROUPS_HEADER: group_names = request.headers.get(WEBUI_AUTH_TRUSTED_GROUPS_HEADER, '').split(',') group_names = [name.strip() for name in group_names if name.strip()] if group_names: - Groups.sync_groups_by_group_names(user.id, group_names, db=db) + await Groups.sync_groups_by_group_names(user.id, group_names, db=db) if WEBUI_AUTH_TRUSTED_ROLE_HEADER: trusted_role = request.headers.get(WEBUI_AUTH_TRUSTED_ROLE_HEADER, '').lower().strip() if trusted_role in {'admin', 'user', 'pending'}: if user.role != trusted_role: - Users.update_user_role_by_id(user.id, trusted_role, db=db) + await Users.update_user_role_by_id(user.id, trusted_role, db=db) elif trusted_role: log.warning(f'Ignoring invalid trusted role header value: {trusted_role}') @@ -594,14 +594,14 @@ async def signin( admin_email = 'admin@localhost' admin_password = 'admin' - if Users.get_user_by_email(admin_email.lower(), db=db): - user = Auths.authenticate_user( + if await Users.get_user_by_email(admin_email.lower(), db=db): + user = await Auths.authenticate_user( admin_email.lower(), lambda pw: verify_password(admin_password, pw), db=db, ) else: - if Users.has_users(db=db): + if await Users.has_users(db=db): raise HTTPException(400, detail=ERROR_MESSAGES.EXISTING_USERS) await signup_handler( @@ -612,7 +612,7 @@ async def signin( db=db, ) - user = Auths.authenticate_user( + user = await Auths.authenticate_user( admin_email.lower(), lambda pw: verify_password(admin_password, pw), db=db, @@ -633,14 +633,14 @@ async def signin( # decode safely — ignore incomplete UTF-8 sequences form_data.password = password_bytes.decode('utf-8', errors='ignore') - user = Auths.authenticate_user( + user = await Auths.authenticate_user( form_data.email.lower(), lambda pw: verify_password(form_data.password, pw), db=db, ) if user: - return create_session_response(request, user, db, response, set_cookie=True) + return await create_session_response(request, user, db, response, set_cookie=True) else: raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED) @@ -657,7 +657,7 @@ async def signup_handler( name: str, profile_image_url: str = '/user.png', *, - db: Session, + db: AsyncSession, ) -> UserModel: """ Core user-creation logic shared by the signup endpoint and @@ -671,7 +671,7 @@ async def signup_handler( # first-user registration can all see an empty table and each get admin. hashed = get_password_hash(password) - user = Auths.insert_new_auth( + user = await Auths.insert_new_auth( email=email.lower(), password=hashed, name=name, @@ -684,9 +684,9 @@ async def signup_handler( # Atomically check if this is the only user *after* the insert. # Only the single user present at this point should become admin. - if Users.get_num_users(db=db) == 1: - Users.update_user_role_by_id(user.id, 'admin', db=db) - user = Users.get_user_by_id(user.id, db=db) + if await Users.get_num_users(db=db) == 1: + await Users.update_user_role_by_id(user.id, 'admin', db=db) + user = await Users.get_user_by_id(user.id, db=db) request.app.state.config.ENABLE_SIGNUP = False if request.app.state.config.WEBHOOK_URL: @@ -701,7 +701,7 @@ async def signup_handler( }, ) - apply_default_group_assignment( + await apply_default_group_assignment( request.app.state.config.DEFAULT_GROUP_ID, user.id, db=db, @@ -715,9 +715,9 @@ async def signup( request: Request, response: Response, form_data: SignupForm, - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - has_users = Users.has_users(db=db) + has_users = await Users.has_users(db=db) if WEBUI_AUTH: if not request.app.state.config.ENABLE_SIGNUP or not request.app.state.config.ENABLE_LOGIN_FORM: @@ -730,7 +730,7 @@ async def signup( if not validate_email_format(form_data.email.lower()): raise HTTPException(status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.INVALID_EMAIL_FORMAT) - if Users.get_user_by_email(form_data.email.lower(), db=db): + if await Users.get_user_by_email(form_data.email.lower(), db=db): raise HTTPException(400, detail=ERROR_MESSAGES.EMAIL_TAKEN) try: @@ -747,7 +747,7 @@ async def signup( form_data.profile_image_url, db=db, ) - return create_session_response(request, user, db, response, set_cookie=True) + return await create_session_response(request, user, db, response, set_cookie=True) except HTTPException: raise except Exception as err: @@ -756,7 +756,7 @@ async def signup( @router.get('/signout') -async def signout(request: Request, response: Response, db: Session = Depends(get_session)): +async def signout(request: Request, response: Response, db: AsyncSession = Depends(get_async_session)): # get auth token from headers or cookies token = None auth_header = request.headers.get('Authorization') @@ -777,7 +777,7 @@ async def signout(request: Request, response: Response, db: Session = Depends(ge if oauth_session_id: response.delete_cookie('oauth_session_id') - session = OAuthSessions.get_session_by_id(oauth_session_id, db=db) + session = await OAuthSessions.get_session_by_id(oauth_session_id, db=db) # If a custom end_session_endpoint is configured (e.g. AWS Cognito), redirect # there directly instead of attempting OIDC discovery. @@ -852,12 +852,12 @@ async def add_user( request: Request, form_data: AddUserForm, user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): if not validate_email_format(form_data.email.lower()): raise HTTPException(status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.INVALID_EMAIL_FORMAT) - if Users.get_user_by_email(form_data.email.lower(), db=db): + if await Users.get_user_by_email(form_data.email.lower(), db=db): raise HTTPException(400, detail=ERROR_MESSAGES.EMAIL_TAKEN) try: @@ -867,7 +867,7 @@ async def add_user( raise HTTPException(400, detail=str(e)) hashed = get_password_hash(form_data.password) - user = Auths.insert_new_auth( + user = await Auths.insert_new_auth( form_data.email.lower(), hashed, form_data.name, @@ -877,7 +877,7 @@ async def add_user( ) if user: - apply_default_group_assignment( + await apply_default_group_assignment( request.app.state.config.DEFAULT_GROUP_ID, user.id, db=db, @@ -909,7 +909,7 @@ async def add_user( @router.get('/admin/details') -async def get_admin_details(request: Request, user=Depends(get_current_user), db: Session = Depends(get_session)): +async def get_admin_details(request: Request, user=Depends(get_current_user), db: AsyncSession = Depends(get_async_session)): if request.app.state.config.SHOW_ADMIN_DETAILS: admin_email = request.app.state.config.ADMIN_EMAIL admin_name = None @@ -917,11 +917,11 @@ async def get_admin_details(request: Request, user=Depends(get_current_user), db log.info(f'Admin details - Email: {admin_email}, Name: {admin_name}') if admin_email: - admin = Users.get_user_by_email(admin_email, db=db) + admin = await Users.get_user_by_email(admin_email, db=db) if admin: admin_name = admin.name else: - admin = Users.get_first_user(db=db) + admin = await Users.get_first_user(db=db) if admin: admin_email = admin.email admin_name = admin.name @@ -1173,10 +1173,10 @@ async def update_ldap_config(request: Request, form_data: LdapConfigForm, user=D # create api key @router.post('/api_key', response_model=ApiKey) -async def generate_api_key(request: Request, user=Depends(get_current_user), db: Session = Depends(get_session)): +async def generate_api_key(request: Request, user=Depends(get_current_user), db: AsyncSession = Depends(get_async_session)): if not request.app.state.config.ENABLE_API_KEYS or ( user.role != 'admin' - and not has_permission(user.id, 'features.api_keys', request.app.state.config.USER_PERMISSIONS) + and not await has_permission(user.id, 'features.api_keys', request.app.state.config.USER_PERMISSIONS) ): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, @@ -1184,7 +1184,7 @@ async def generate_api_key(request: Request, user=Depends(get_current_user), db: ) api_key = create_api_key() - success = Users.update_user_api_key_by_id(user.id, api_key, db=db) + success = await Users.update_user_api_key_by_id(user.id, api_key, db=db) if success: return { @@ -1196,14 +1196,14 @@ async def generate_api_key(request: Request, user=Depends(get_current_user), db: # delete api key @router.delete('/api_key', response_model=bool) -async def delete_api_key(user=Depends(get_current_user), db: Session = Depends(get_session)): - return Users.delete_user_api_key_by_id(user.id, db=db) +async def delete_api_key(user=Depends(get_current_user), db: AsyncSession = Depends(get_async_session)): + return await Users.delete_user_api_key_by_id(user.id, db=db) # get api key @router.get('/api_key', response_model=ApiKey) -async def get_api_key(user=Depends(get_current_user), db: Session = Depends(get_session)): - api_key = Users.get_user_api_key_by_id(user.id, db=db) +async def get_api_key(user=Depends(get_current_user), db: AsyncSession = Depends(get_async_session)): + api_key = await Users.get_user_api_key_by_id(user.id, db=db) if api_key: return { 'api_key': api_key, @@ -1227,7 +1227,7 @@ async def token_exchange( response: Response, provider: str, form_data: TokenExchangeForm, - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """ Exchange an external OAuth provider token for an OpenWebUI JWT. @@ -1296,14 +1296,14 @@ async def token_exchange( email = email.lower() # Try to find the user by OAuth sub - user = Users.get_user_by_oauth_sub(provider, sub, db=db) + user = await Users.get_user_by_oauth_sub(provider, sub, db=db) if not user and OAUTH_MERGE_ACCOUNTS_BY_EMAIL.value: # Try to find by email if merge is enabled - user = Users.get_user_by_email(email, db=db) + user = await Users.get_user_by_email(email, db=db) if user: # Link the OAuth sub to this user - Users.update_user_oauth_by_id(user.id, provider, sub, db=db) + await Users.update_user_oauth_by_id(user.id, provider, sub, db=db) if not user: raise HTTPException( @@ -1311,4 +1311,4 @@ async def token_exchange( detail='User not found. Please sign in via the web interface first.', ) - return create_session_response(request, user, db) + return await create_session_response(request, user, db) diff --git a/backend/open_webui/routers/automations.py b/backend/open_webui/routers/automations.py index 0f85115720..504bc726d2 100644 --- a/backend/open_webui/routers/automations.py +++ b/backend/open_webui/routers/automations.py @@ -3,7 +3,7 @@ import logging from typing import Optional from fastapi import APIRouter, Depends, HTTPException, Request, status -from sqlalchemy.orm import Session +from sqlalchemy.ext.asyncio import AsyncSession from open_webui.models.automations import ( Automations, @@ -23,7 +23,7 @@ from open_webui.utils.automations import ( ) from open_webui.utils.auth import get_verified_user, get_admin_user from open_webui.utils.access_control import has_permission -from open_webui.internal.db import get_session +from open_webui.internal.db import get_async_session from open_webui.constants import ERROR_MESSAGES log = logging.getLogger(__name__) @@ -38,8 +38,8 @@ PAGE_ITEM_COUNT = 30 ############################ -def check_automations_permission(request, user): - if user.role != 'admin' and not has_permission( +async def check_automations_permission(request, user): + if user.role != 'admin' and not await has_permission( user.id, 'features.automations', request.app.state.config.USER_PERMISSIONS ): raise HTTPException( @@ -61,7 +61,7 @@ def check_automation_access(automation, user): ) -def check_automation_limits(request, user, rrule_str: str, db, is_create: bool = False): +async def check_automation_limits(request, user, rrule_str: str, db, is_create: bool = False): """Enforce global automation limits. Admins bypass all checks.""" if user.role == 'admin': return @@ -71,7 +71,7 @@ def check_automation_limits(request, user, rrule_str: str, db, is_create: bool = max_count = request.app.state.config.AUTOMATION_MAX_COUNT if max_count: max_count = int(max_count) - if max_count > 0 and Automations.count_by_user(user.id, db=db) >= max_count: + if max_count > 0 and await Automations.count_by_user(user.id, db=db) >= max_count: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=f'Automation limit reached ({max_count})', @@ -90,9 +90,9 @@ def check_automation_limits(request, user, rrule_str: str, db, is_create: bool = ) -def enrich_automation(automation: AutomationModel, db: Session, tz: str = None) -> AutomationResponse: +async def enrich_automation(automation: AutomationModel, db: AsyncSession, tz: str = None) -> AutomationResponse: """Full enrichment for single-item views (includes next_runs computation).""" - last_run = AutomationRuns.get_latest(automation.id, db=db) + last_run = await AutomationRuns.get_latest(automation.id, db=db) return AutomationResponse( **automation.model_dump(), last_run=last_run, @@ -112,14 +112,14 @@ async def get_automation_items( status: Optional[str] = None, page: Optional[int] = 1, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_automations_permission(request, user) + await check_automations_permission(request, user) limit = PAGE_ITEM_COUNT page = max(1, page) skip = (page - 1) * limit - result = Automations.search_automations( + result = await Automations.search_automations( user_id=user.id, query=query, status=status, @@ -130,7 +130,7 @@ async def get_automation_items( # Batch-fetch latest runs in a single query instead of N+1 ids = [item.id for item in result.items] - latest_runs = AutomationRuns.get_latest_batch(ids, db=db) if ids else {} + latest_runs = await AutomationRuns.get_latest_batch(ids, db=db) if ids else {} return { 'items': [ @@ -154,9 +154,9 @@ async def create_new_automation( request: Request, form_data: AutomationForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_automations_permission(request, user) + await check_automations_permission(request, user) try: validate_rrule(form_data.data.rrule) except ValueError as e: @@ -165,7 +165,7 @@ async def create_new_automation( detail=str(e), ) - check_automation_limits(request, user, form_data.data.rrule, db, is_create=True) + await check_automation_limits(request, user, form_data.data.rrule, db, is_create=True) # Validate terminal server exists if linked if form_data.data.terminal and form_data.data.terminal.server_id: @@ -177,8 +177,8 @@ async def create_new_automation( ) tz = user.timezone - automation = Automations.insert(user.id, form_data, next_run_ns(form_data.data.rrule, tz=tz), db=db) - return enrich_automation(automation, db, tz=tz) + automation = await Automations.insert(user.id, form_data, next_run_ns(form_data.data.rrule, tz=tz), db=db) + return await enrich_automation(automation, db, tz=tz) ############################ @@ -191,12 +191,12 @@ async def get_automation_by_id( request: Request, id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_automations_permission(request, user) - automation = Automations.get_by_id(id, db=db) + await check_automations_permission(request, user) + automation = await Automations.get_by_id(id, db=db) check_automation_access(automation, user) - return enrich_automation(automation, db, tz=user.timezone) + return await enrich_automation(automation, db, tz=user.timezone) ############################ @@ -210,10 +210,10 @@ async def update_automation_by_id( id: str, form_data: AutomationForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_automations_permission(request, user) - automation = Automations.get_by_id(id, db=db) + await check_automations_permission(request, user) + automation = await Automations.get_by_id(id, db=db) check_automation_access(automation, user) try: @@ -224,7 +224,7 @@ async def update_automation_by_id( detail=str(e), ) - check_automation_limits(request, user, form_data.data.rrule, db, is_create=False) + await check_automation_limits(request, user, form_data.data.rrule, db, is_create=False) # Validate terminal server exists if linked if form_data.data.terminal and form_data.data.terminal.server_id: @@ -236,8 +236,8 @@ async def update_automation_by_id( ) tz = user.timezone - updated = Automations.update_by_id(id, form_data, next_run_ns(form_data.data.rrule, tz=tz), db=db) - return enrich_automation(updated, db, tz=tz) + updated = await Automations.update_by_id(id, form_data, next_run_ns(form_data.data.rrule, tz=tz), db=db) + return await enrich_automation(updated, db, tz=tz) ############################ @@ -250,13 +250,13 @@ async def toggle_automation_by_id( request: Request, id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_automations_permission(request, user) - automation = Automations.get_by_id(id, db=db) + await check_automations_permission(request, user) + automation = await Automations.get_by_id(id, db=db) check_automation_access(automation, user) - toggled = Automations.toggle(id, next_run_ns(automation.data['rrule'], tz=user.timezone), db=db) - return enrich_automation(toggled, db, tz=user.timezone) + toggled = await Automations.toggle(id, next_run_ns(automation.data['rrule'], tz=user.timezone), db=db) + return await enrich_automation(toggled, db, tz=user.timezone) ############################ @@ -269,13 +269,13 @@ async def run_automation_by_id( request: Request, id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_automations_permission(request, user) - automation = Automations.get_by_id(id, db=db) + await check_automations_permission(request, user) + automation = await Automations.get_by_id(id, db=db) check_automation_access(automation, user) asyncio.create_task(execute_automation(request.app, automation)) - return enrich_automation(automation, db, tz=user.timezone) + return await enrich_automation(automation, db, tz=user.timezone) ############################ @@ -288,13 +288,13 @@ async def delete_automation_by_id( request: Request, id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_automations_permission(request, user) - automation = Automations.get_by_id(id, db=db) + await check_automations_permission(request, user) + automation = await Automations.get_by_id(id, db=db) check_automation_access(automation, user) - AutomationRuns.delete_by_automation(id, db=db) - return Automations.delete(id, db=db) + await AutomationRuns.delete_by_automation(id, db=db) + return await Automations.delete(id, db=db) ############################ @@ -309,9 +309,9 @@ async def get_automation_runs( skip: int = 0, limit: int = 50, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_automations_permission(request, user) - automation = Automations.get_by_id(id, db=db) + await check_automations_permission(request, user) + automation = await Automations.get_by_id(id, db=db) check_automation_access(automation, user) - return AutomationRuns.get_by_automation(id, skip=skip, limit=limit, db=db) + return await AutomationRuns.get_by_automation(id, skip=skip, limit=limit, db=db) diff --git a/backend/open_webui/routers/channels.py b/backend/open_webui/routers/channels.py index 71bb34394f..6610ee2eca 100644 --- a/backend/open_webui/routers/channels.py +++ b/backend/open_webui/routers/channels.py @@ -64,22 +64,22 @@ from open_webui.utils.auth import get_admin_user, get_verified_user from open_webui.utils.access_control import has_permission from open_webui.utils.webhook import post_webhook from open_webui.utils.channels import extract_mentions, replace_mentions -from open_webui.internal.db import get_session -from sqlalchemy.orm import Session +from open_webui.internal.db import get_async_session +from sqlalchemy.ext.asyncio import AsyncSession log = logging.getLogger(__name__) router = APIRouter() -def channel_has_access( +async def channel_has_access( user_id: str, channel: ChannelModel, permission: str = 'read', strict: bool = True, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ) -> bool: - if AccessGrants.has_access( + if await AccessGrants.has_access( user_id=user_id, resource_type='channel', resource_id=channel.id, @@ -94,8 +94,8 @@ def channel_has_access( return False -def get_channel_users_with_access(channel: ChannelModel, permission: str = 'read', db: Optional[Session] = None): - return AccessGrants.get_users_with_access( +async def get_channel_users_with_access(channel: ChannelModel, permission: str = 'read', db: Optional[AsyncSession] = None): + return await AccessGrants.get_users_with_access( resource_type='channel', resource_id=channel.id, permission=permission, @@ -133,7 +133,7 @@ def get_channel_permitted_group_and_user_ids( ############################ -def check_channels_access(request: Request, user: Optional[UserModel] = None): +async def check_channels_access(request: Request, user: Optional[UserModel] = None): """Dependency to ensure channels are globally enabled.""" if not request.app.state.config.ENABLE_CHANNELS: raise HTTPException( @@ -142,7 +142,7 @@ def check_channels_access(request: Request, user: Optional[UserModel] = None): ) if user: - if user.role != 'admin' and not has_permission( + if user.role != 'admin' and not await has_permission( user.id, 'features.channels', request.app.state.config.USER_PERMISSIONS ): raise HTTPException( @@ -168,19 +168,19 @@ class ChannelListItemResponse(ChannelModel): async def get_channels( request: Request, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) + await check_channels_access(request, user) - channels = Channels.get_channels_by_user_id(user.id, db=db) + channels = await Channels.get_channels_by_user_id(user.id, db=db) channel_list = [] for channel in channels: - last_message = Messages.get_last_message_by_channel_id(channel.id, db=db) + last_message = await Messages.get_last_message_by_channel_id(channel.id, db=db) last_message_at = last_message.created_at if last_message else None - channel_member = Channels.get_member_by_channel_and_user_id(channel.id, user.id, db=db) + channel_member = await Channels.get_member_by_channel_and_user_id(channel.id, user.id, db=db) unread_count = ( - Messages.get_unread_message_count(channel.id, user.id, channel_member.last_read_at, db=db) + await Messages.get_unread_message_count(channel.id, user.id, channel_member.last_read_at, db=db) if channel_member else 0 ) @@ -188,15 +188,15 @@ async def get_channels( user_ids = None users = None if channel.type == 'dm': - user_ids = [member.user_id for member in Channels.get_members_by_channel_id(channel.id, db=db)] + user_ids = [member.user_id for member in await Channels.get_members_by_channel_id(channel.id, db=db)] users = [ UserIdNameStatusResponse( **{ - **user.model_dump(), - 'is_active': Users.is_active(user), + **u.model_dump(), + 'is_active': Users.is_active(u), } ) - for user in Users.get_users_by_user_ids(user_ids, db=db) + for u in await Users.get_users_by_user_ids(user_ids, db=db) ] channel_list.append( @@ -216,12 +216,12 @@ async def get_channels( async def get_all_channels( request: Request, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) + await check_channels_access(request, user) if user.role == 'admin': - return Channels.get_channels(db=db) - return Channels.get_channels_by_user_id(user.id, db=db) + return await Channels.get_channels(db=db) + return await Channels.get_channels_by_user_id(user.id, db=db) ############################ @@ -234,14 +234,14 @@ async def get_dm_channel_by_user_id( request: Request, user_id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) + await check_channels_access(request, user) try: - existing_channel = Channels.get_dm_channel_by_user_ids([user.id, user_id], db=db) + existing_channel = await Channels.get_dm_channel_by_user_ids([user.id, user_id], db=db) if existing_channel: participant_ids = [ - member.user_id for member in Channels.get_members_by_channel_id(existing_channel.id, db=db) + member.user_id for member in await Channels.get_members_by_channel_id(existing_channel.id, db=db) ] await emit_to_users( @@ -251,10 +251,10 @@ async def get_dm_channel_by_user_id( ) await enter_room_for_users(f'channel:{existing_channel.id}', participant_ids) - Channels.update_member_active_status(existing_channel.id, user.id, True, db=db) + await Channels.update_member_active_status(existing_channel.id, user.id, True, db=db) return ChannelModel(**existing_channel.model_dump()) - channel = Channels.insert_new_channel( + channel = await Channels.insert_new_channel( CreateChannelForm( type='dm', name='', @@ -265,7 +265,7 @@ async def get_dm_channel_by_user_id( ) if channel: - participant_ids = [member.user_id for member in Channels.get_members_by_channel_id(channel.id, db=db)] + participant_ids = [member.user_id for member in await Channels.get_members_by_channel_id(channel.id, db=db)] await emit_to_users( 'events:channel', @@ -292,9 +292,9 @@ async def create_new_channel( request: Request, form_data: CreateChannelForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) + await check_channels_access(request, user) if form_data.type not in ['group', 'dm'] and user.role != 'admin': # Only admins can create standard channels (joined by default) @@ -305,10 +305,10 @@ async def create_new_channel( try: if form_data.type == 'dm': - existing_channel = Channels.get_dm_channel_by_user_ids([user.id, *form_data.user_ids], db=db) + existing_channel = await Channels.get_dm_channel_by_user_ids([user.id, *form_data.user_ids], db=db) if existing_channel: participant_ids = [ - member.user_id for member in Channels.get_members_by_channel_id(existing_channel.id, db=db) + member.user_id for member in await Channels.get_members_by_channel_id(existing_channel.id, db=db) ] await emit_to_users( 'events:channel', @@ -317,13 +317,13 @@ async def create_new_channel( ) await enter_room_for_users(f'channel:{existing_channel.id}', participant_ids) - Channels.update_member_active_status(existing_channel.id, user.id, True, db=db) + await Channels.update_member_active_status(existing_channel.id, user.id, True, db=db) return ChannelModel(**existing_channel.model_dump()) - channel = Channels.insert_new_channel(form_data, user.id, db=db) + channel = await Channels.insert_new_channel(form_data, user.id, db=db) if channel: - participant_ids = [member.user_id for member in Channels.get_members_by_channel_id(channel.id, db=db)] + participant_ids = [member.user_id for member in await Channels.get_members_by_channel_id(channel.id, db=db)] await emit_to_users( 'events:channel', @@ -358,10 +358,10 @@ async def get_channel_by_id( request: Request, id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) - channel = Channels.get_channel_by_id(id, db=db) + await check_channels_access(request, user) + channel = await Channels.get_channel_by_id(id, db=db) if not channel: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) @@ -369,23 +369,23 @@ async def get_channel_by_id( users = None if channel.type in ['group', 'dm']: - if not Channels.is_user_channel_member(channel.id, user.id, db=db): + if not await Channels.is_user_channel_member(channel.id, user.id, db=db): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) - user_ids = [member.user_id for member in Channels.get_members_by_channel_id(channel.id, db=db)] + user_ids = [member.user_id for member in await Channels.get_members_by_channel_id(channel.id, db=db)] users = [ UserIdNameStatusResponse( **{ - **user.model_dump(), - 'is_active': Users.is_active(user), + **u.model_dump(), + 'is_active': Users.is_active(u), } ) - for user in Users.get_users_by_user_ids(user_ids, db=db) + for u in await Users.get_users_by_user_ids(user_ids, db=db) ] - channel_member = Channels.get_member_by_channel_and_user_id(channel.id, user.id, db=db) - unread_count = Messages.get_unread_message_count( + channel_member = await Channels.get_member_by_channel_and_user_id(channel.id, user.id, db=db) + unread_count = await Messages.get_unread_message_count( channel.id, user.id, channel_member.last_read_at if channel_member else None ) @@ -394,7 +394,7 @@ async def get_channel_by_id( **channel.model_dump(), 'user_ids': user_ids, 'users': users, - 'is_manager': Channels.is_user_channel_manager(channel.id, user.id, db=db), + 'is_manager': await Channels.is_user_channel_manager(channel.id, user.id, db=db), 'write_access': True, 'user_count': len(user_ids), 'last_read_at': channel_member.last_read_at if channel_member else None, @@ -402,10 +402,10 @@ async def get_channel_by_id( } ) else: - if user.role != 'admin' and not channel_has_access(user.id, channel, permission='read', db=db): + if user.role != 'admin' and not await channel_has_access(user.id, channel, permission='read', db=db): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) - write_access = channel_has_access( + write_access = await channel_has_access( user.id, channel, permission='write', @@ -413,10 +413,10 @@ async def get_channel_by_id( db=db, ) - user_count = len(get_channel_users_with_access(channel, 'read', db=db)) + user_count = len(await get_channel_users_with_access(channel, 'read', db=db)) - channel_member = Channels.get_member_by_channel_and_user_id(channel.id, user.id, db=db) - unread_count = Messages.get_unread_message_count( + channel_member = await Channels.get_member_by_channel_and_user_id(channel.id, user.id, db=db) + unread_count = await Messages.get_unread_message_count( channel.id, user.id, channel_member.last_read_at if channel_member else None ) @@ -425,7 +425,7 @@ async def get_channel_by_id( **channel.model_dump(), 'user_ids': user_ids, 'users': users, - 'is_manager': Channels.is_user_channel_manager(channel.id, user.id, db=db), + 'is_manager': await Channels.is_user_channel_manager(channel.id, user.id, db=db), 'write_access': write_access or user.role == 'admin', 'user_count': user_count, 'last_read_at': channel_member.last_read_at if channel_member else None, @@ -451,11 +451,11 @@ async def get_channel_members_by_id( direction: Optional[str] = None, page: Optional[int] = 1, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) + await check_channels_access(request, user) - channel = Channels.get_channel_by_id(id, db=db) + channel = await Channels.get_channel_by_id(id, db=db) if not channel: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) @@ -465,19 +465,19 @@ async def get_channel_members_by_id( skip = (page - 1) * limit if channel.type in ['group', 'dm']: - if not Channels.is_user_channel_member(channel.id, user.id, db=db): + if not await Channels.is_user_channel_member(channel.id, user.id, db=db): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) else: - if user.role != 'admin' and not channel_has_access(user.id, channel, permission='read', db=db): + if user.role != 'admin' and not await channel_has_access(user.id, channel, permission='read', db=db): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) if channel.type == 'dm': - user_ids = [member.user_id for member in Channels.get_members_by_channel_id(channel.id, db=db)] - users = Users.get_users_by_user_ids(user_ids, db=db) - total = len(users) + user_ids = [member.user_id for member in await Channels.get_members_by_channel_id(channel.id, db=db)] + fetched_users = await Users.get_users_by_user_ids(user_ids, db=db) + total = len(fetched_users) return { - 'users': [UserModelResponse(**user.model_dump(), is_active=Users.is_active(user)) for user in users], + 'users': [UserModelResponse(**u.model_dump(), is_active=Users.is_active(u)) for u in fetched_users], 'total': total, } else: @@ -499,13 +499,13 @@ async def get_channel_members_by_id( filter['user_ids'] = permitted_ids.get('user_ids') filter['group_ids'] = permitted_ids.get('group_ids') - result = Users.get_users(filter=filter, skip=skip, limit=limit, db=db) + result = await Users.get_users(filter=filter, skip=skip, limit=limit, db=db) - users = result['users'] + fetched_users = result['users'] total = result['total'] return { - 'users': [UserModelResponse(**user.model_dump(), is_active=Users.is_active(user)) for user in users], + 'users': [UserModelResponse(**u.model_dump(), is_active=Users.is_active(u)) for u in fetched_users], 'total': total, } @@ -525,17 +525,17 @@ async def update_is_active_member_by_id_and_user_id( id: str, form_data: UpdateActiveMemberForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) - channel = Channels.get_channel_by_id(id, db=db) + await check_channels_access(request, user) + channel = await Channels.get_channel_by_id(id, db=db) if not channel: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) - if not Channels.is_user_channel_member(channel.id, user.id, db=db): + if not await Channels.is_user_channel_member(channel.id, user.id, db=db): raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) - Channels.update_member_active_status(channel.id, user.id, form_data.is_active, db=db) + await Channels.update_member_active_status(channel.id, user.id, form_data.is_active, db=db) return True @@ -555,10 +555,10 @@ async def add_members_by_id( id: str, form_data: UpdateMembersForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) - channel = Channels.get_channel_by_id(id, db=db) + await check_channels_access(request, user) + channel = await Channels.get_channel_by_id(id, db=db) if not channel: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) @@ -566,7 +566,7 @@ async def add_members_by_id( raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) try: - memberships = Channels.add_members_to_channel( + memberships = await Channels.add_members_to_channel( channel.id, user.id, form_data.user_ids, form_data.group_ids, db=db ) @@ -591,11 +591,11 @@ async def remove_members_by_id( id: str, form_data: RemoveMembersForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) + await check_channels_access(request, user) - channel = Channels.get_channel_by_id(id, db=db) + channel = await Channels.get_channel_by_id(id, db=db) if not channel: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) @@ -603,7 +603,7 @@ async def remove_members_by_id( raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) try: - deleted = Channels.remove_members_from_channel(channel.id, form_data.user_ids, db=db) + deleted = await Channels.remove_members_from_channel(channel.id, form_data.user_ids, db=db) return deleted except Exception as e: @@ -622,11 +622,11 @@ async def update_channel_by_id( id: str, form_data: ChannelForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) + await check_channels_access(request, user) - channel = Channels.get_channel_by_id(id, db=db) + channel = await Channels.get_channel_by_id(id, db=db) if not channel: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) @@ -634,7 +634,7 @@ async def update_channel_by_id( raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) try: - channel = Channels.update_channel_by_id(id, form_data, db=db) + channel = await Channels.update_channel_by_id(id, form_data, db=db) return ChannelModel(**channel.model_dump()) except Exception as e: log.exception(e) @@ -651,11 +651,11 @@ async def delete_channel_by_id( request: Request, id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) + await check_channels_access(request, user) - channel = Channels.get_channel_by_id(id, db=db) + channel = await Channels.get_channel_by_id(id, db=db) if not channel: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) @@ -663,7 +663,7 @@ async def delete_channel_by_id( raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) try: - Channels.delete_channel_by_id(id, db=db) + await Channels.delete_channel_by_id(id, db=db) return True except Exception as e: log.exception(e) @@ -695,40 +695,40 @@ async def get_channel_messages( skip: int = 0, limit: int = 50, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) - channel = Channels.get_channel_by_id(id, db=db) + await check_channels_access(request, user) + channel = await Channels.get_channel_by_id(id, db=db) if not channel: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) if channel.type in ['group', 'dm']: - if not Channels.is_user_channel_member(channel.id, user.id, db=db): + if not await Channels.is_user_channel_member(channel.id, user.id, db=db): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) else: - if user.role != 'admin' and not channel_has_access(user.id, channel, permission='read', db=db): + if user.role != 'admin' and not await channel_has_access(user.id, channel, permission='read', db=db): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) - channel_member = Channels.join_channel(id, user.id, db=db) # Ensure user is a member of the channel + channel_member = await Channels.join_channel(id, user.id, db=db) # Ensure user is a member of the channel - message_list = Messages.get_messages_by_channel_id(id, skip, limit, db=db) + message_list = await Messages.get_messages_by_channel_id(id, skip, limit, db=db) if not message_list: return [] # Batch fetch all users in a single query (fixes N+1 problem) user_ids = list(set(m.user_id for m in message_list)) - users = {u.id: u for u in Users.get_users_by_user_ids(user_ids, db=db)} + fetched_users = {u.id: u for u in await Users.get_users_by_user_ids(user_ids, db=db)} messages = [] for message in message_list: - thread_replies = Messages.get_thread_replies_by_message_id(message.id, db=db) + thread_replies = await Messages.get_thread_replies_by_message_id(message.id, db=db) latest_thread_reply_at = thread_replies[0].created_at if thread_replies else None # Use message.user if present (for webhooks), otherwise look up by user_id user_info = message.user - if user_info is None and message.user_id in users: - user_info = UserNameResponse(**users[message.user_id].model_dump()) + if user_info is None and message.user_id in fetched_users: + user_info = UserNameResponse(**fetched_users[message.user_id].model_dump()) messages.append( MessageUserResponse( @@ -736,7 +736,7 @@ async def get_channel_messages( **message.model_dump(), 'reply_count': len(thread_replies), 'latest_reply_at': latest_thread_reply_at, - 'reactions': Messages.get_reactions_by_message_id(message.id, db=db), + 'reactions': await Messages.get_reactions_by_message_id(message.id, db=db), 'user': user_info, } ) @@ -758,32 +758,32 @@ async def get_pinned_channel_messages( id: str, page: int = 1, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) - channel = Channels.get_channel_by_id(id, db=db) + await check_channels_access(request, user) + channel = await Channels.get_channel_by_id(id, db=db) if not channel: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) if channel.type in ['group', 'dm']: - if not Channels.is_user_channel_member(channel.id, user.id, db=db): + if not await Channels.is_user_channel_member(channel.id, user.id, db=db): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) else: - if user.role != 'admin' and not channel_has_access(user.id, channel, permission='read', db=db): + if user.role != 'admin' and not await channel_has_access(user.id, channel, permission='read', db=db): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) page = max(1, page) skip = (page - 1) * PAGE_ITEM_COUNT_PINNED limit = PAGE_ITEM_COUNT_PINNED - message_list = Messages.get_pinned_messages_by_channel_id(id, skip, limit, db=db) + message_list = await Messages.get_pinned_messages_by_channel_id(id, skip, limit, db=db) if not message_list: return [] # Batch fetch all users in a single query (fixes N+1 problem) user_ids = list(set(m.user_id for m in message_list)) - users = {u.id: u for u in Users.get_users_by_user_ids(user_ids, db=db)} + fetched_users = {u.id: u for u in await Users.get_users_by_user_ids(user_ids, db=db)} messages = [] for message in message_list: @@ -795,8 +795,8 @@ async def get_pinned_channel_messages( name=webhook_info.get('name') or 'Webhook', role='webhook', ) - elif message.user_id in users: - user_info = UserNameResponse(**users[message.user_id].model_dump()) + elif message.user_id in fetched_users: + user_info = UserNameResponse(**fetched_users[message.user_id].model_dump()) else: user_info = None @@ -804,7 +804,7 @@ async def get_pinned_channel_messages( MessageWithReactionsResponse( **{ **message.model_dump(), - 'reactions': Messages.get_reactions_by_message_id(message.id, db=db), + 'reactions': await Messages.get_reactions_by_message_id(message.id, db=db), 'user': user_info, } ) @@ -823,12 +823,12 @@ async def send_notification(request, channel, message, active_user_ids, db=None) webui_url = request.app.state.config.WEBUI_URL enable_user_webhooks = request.app.state.config.ENABLE_USER_WEBHOOKS - users = get_channel_users_with_access(channel, 'read', db=db) + users = await get_channel_users_with_access(channel, 'read', db=db) - for user in users: - if (user.id not in active_user_ids) and Channels.is_user_channel_member(channel.id, user.id, db=db): - if enable_user_webhooks and user.settings: - webhook_url = user.settings.ui.get('notifications', {}).get('webhook_url', None) + for u in users: + if (u.id not in active_user_ids) and await Channels.is_user_channel_member(channel.id, u.id, db=db): + if enable_user_webhooks and u.settings: + webhook_url = u.settings.ui.get('notifications', {}).get('webhook_url', None) if webhook_url: await post_webhook( name, @@ -846,7 +846,7 @@ async def send_notification(request, channel, message, active_user_ids, db=None) async def model_response_handler(request, channel, message, user, db=None): - MODELS = {model['id']: model for model in get_filtered_models(await get_all_models(request, user=user), user)} + MODELS = {model['id']: model for model in await get_filtered_models(await get_all_models(request, user=user), user)} mentions = extract_mentions(message.content) message_content = replace_mentions(message.content) @@ -877,11 +877,11 @@ async def model_response_handler(request, channel, message, user, db=None): if model: try: # reverse to get in chronological order - thread_messages = Messages.get_messages_by_parent_id( + thread_messages = (await Messages.get_messages_by_parent_id( channel.id, message.parent_id if message.parent_id else message.id, db=db, - )[::-1] + ))[::-1] response_message, channel = await new_message_handler( request, @@ -908,7 +908,7 @@ async def model_response_handler(request, channel, message, user, db=None): for thread_message in thread_messages: message_user = None if thread_message.user_id not in message_users: - message_user = Users.get_user_by_id(thread_message.user_id, db=db) + message_user = await Users.get_user_by_id(thread_message.user_id, db=db) message_users[thread_message.user_id] = message_user else: message_user = message_users[thread_message.user_id] @@ -928,7 +928,7 @@ async def model_response_handler(request, channel, message, user, db=None): if file.get('type', '') == 'image': images.append(file.get('url', '')) elif file.get('content_type', '').startswith('image/'): - image = get_image_base64_from_file_id(file.get('id', '')) + image = await get_image_base64_from_file_id(file.get('id', '')) if image: images.append(image) @@ -1017,15 +1017,15 @@ async def model_response_handler(request, channel, message, user, db=None): async def new_message_handler(request: Request, id: str, form_data: MessageForm, user, db): - channel = Channels.get_channel_by_id(id, db=db) + channel = await Channels.get_channel_by_id(id, db=db) if not channel: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) if channel.type in ['group', 'dm']: - if not Channels.is_user_channel_member(channel.id, user.id, db=db): + if not await Channels.is_user_channel_member(channel.id, user.id, db=db): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) else: - if user.role != 'admin' and not channel_has_access( + if user.role != 'admin' and not await channel_has_access( user.id, channel, permission='write', @@ -1035,15 +1035,15 @@ async def new_message_handler(request: Request, id: str, form_data: MessageForm, raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) try: - message = Messages.insert_new_message(form_data, channel.id, user.id, db=db) + message = await Messages.insert_new_message(form_data, channel.id, user.id, db=db) if message: if channel.type in ['group', 'dm']: - members = Channels.get_members_by_channel_id(channel.id, db=db) + members = await Channels.get_members_by_channel_id(channel.id, db=db) for member in members: if not member.is_active: - Channels.update_member_active_status(channel.id, member.user_id, True, db=db) + await Channels.update_member_active_status(channel.id, member.user_id, True, db=db) - message = Messages.get_message_by_id(message.id, db=db) + message = await Messages.get_message_by_id(message.id, db=db) event_data = { 'channel_id': channel.id, 'message_id': message.id, @@ -1063,7 +1063,7 @@ async def new_message_handler(request: Request, id: str, form_data: MessageForm, if message.parent_id: # If this message is a reply, emit to the parent message as well - parent_message = Messages.get_message_by_id(message.parent_id, db=db) + parent_message = await Messages.get_message_by_id(message.parent_id, db=db) if parent_message: await sio.emit( @@ -1095,16 +1095,16 @@ async def post_new_message( form_data: MessageForm, background_tasks: BackgroundTasks, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) + await check_channels_access(request, user) try: message, channel = await new_message_handler(request, id, form_data, user, db) try: if files := message.data.get('files', []): for file in files: - Channels.set_file_message_id_in_channel_by_id(channel.id, file.get('id', ''), message.id, db=db) + await Channels.set_file_message_id_in_channel_by_id(channel.id, file.get('id', ''), message.id, db=db) except Exception as e: log.debug(e) @@ -1144,28 +1144,28 @@ async def get_channel_message( id: str, message_id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) - channel = Channels.get_channel_by_id(id, db=db) + await check_channels_access(request, user) + channel = await Channels.get_channel_by_id(id, db=db) if not channel: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) if channel.type in ['group', 'dm']: - if not Channels.is_user_channel_member(channel.id, user.id, db=db): + if not await Channels.is_user_channel_member(channel.id, user.id, db=db): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) else: - if user.role != 'admin' and not channel_has_access(user.id, channel, permission='read', db=db): + if user.role != 'admin' and not await channel_has_access(user.id, channel, permission='read', db=db): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) - message = Messages.get_message_by_id(message_id, db=db) + message = await Messages.get_message_by_id(message_id, db=db) if not message: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) if message.channel_id != id: raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.DEFAULT()) - message_user = Users.get_user_by_id(message.user_id, db=db) + message_user = await Users.get_user_by_id(message.user_id, db=db) return MessageResponse( **{ **message.model_dump(), @@ -1185,21 +1185,21 @@ async def get_channel_message_data( id: str, message_id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) - channel = Channels.get_channel_by_id(id, db=db) + await check_channels_access(request, user) + channel = await Channels.get_channel_by_id(id, db=db) if not channel: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) if channel.type in ['group', 'dm']: - if not Channels.is_user_channel_member(channel.id, user.id, db=db): + if not await Channels.is_user_channel_member(channel.id, user.id, db=db): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) else: - if user.role != 'admin' and not channel_has_access(user.id, channel, permission='read', db=db): + if user.role != 'admin' and not await channel_has_access(user.id, channel, permission='read', db=db): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) - message = Messages.get_message_by_id(message_id, db=db) + message = await Messages.get_message_by_id(message_id, db=db) if not message: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) @@ -1225,21 +1225,21 @@ async def pin_channel_message( message_id: str, form_data: PinMessageForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) - channel = Channels.get_channel_by_id(id, db=db) + await check_channels_access(request, user) + channel = await Channels.get_channel_by_id(id, db=db) if not channel: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) if channel.type in ['group', 'dm']: - if not Channels.is_user_channel_member(channel.id, user.id, db=db): + if not await Channels.is_user_channel_member(channel.id, user.id, db=db): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) else: - if user.role != 'admin' and not channel_has_access(user.id, channel, permission='read', db=db): + if user.role != 'admin' and not await channel_has_access(user.id, channel, permission='read', db=db): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) - message = Messages.get_message_by_id(message_id, db=db) + message = await Messages.get_message_by_id(message_id, db=db) if not message: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) @@ -1247,9 +1247,9 @@ async def pin_channel_message( raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.DEFAULT()) try: - Messages.update_is_pinned_by_id(message_id, form_data.is_pinned, user.id, db=db) - message = Messages.get_message_by_id(message_id, db=db) - message_user = Users.get_user_by_id(message.user_id, db=db) + await Messages.update_is_pinned_by_id(message_id, form_data.is_pinned, user.id, db=db) + message = await Messages.get_message_by_id(message_id, db=db) + message_user = await Users.get_user_by_id(message.user_id, db=db) return MessageUserResponse( **{ **message.model_dump(), @@ -1274,35 +1274,35 @@ async def get_channel_thread_messages( skip: int = 0, limit: int = 50, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) - channel = Channels.get_channel_by_id(id, db=db) + await check_channels_access(request, user) + channel = await Channels.get_channel_by_id(id, db=db) if not channel: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) if channel.type in ['group', 'dm']: - if not Channels.is_user_channel_member(channel.id, user.id, db=db): + if not await Channels.is_user_channel_member(channel.id, user.id, db=db): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) else: - if user.role != 'admin' and not channel_has_access(user.id, channel, permission='read', db=db): + if user.role != 'admin' and not await channel_has_access(user.id, channel, permission='read', db=db): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) - message_list = Messages.get_messages_by_parent_id(id, message_id, skip, limit, db=db) + message_list = await Messages.get_messages_by_parent_id(id, message_id, skip, limit, db=db) if not message_list: return [] # Batch fetch all users in a single query (fixes N+1 problem) user_ids = list(set(m.user_id for m in message_list)) - users = {u.id: u for u in Users.get_users_by_user_ids(user_ids, db=db)} + fetched_users = {u.id: u for u in await Users.get_users_by_user_ids(user_ids, db=db)} messages = [] for message in message_list: # Use message.user if present (for webhooks), otherwise look up by user_id user_info = message.user - if user_info is None and message.user_id in users: - user_info = UserNameResponse(**users[message.user_id].model_dump()) + if user_info is None and message.user_id in fetched_users: + user_info = UserNameResponse(**fetched_users[message.user_id].model_dump()) messages.append( MessageUserResponse( @@ -1310,7 +1310,7 @@ async def get_channel_thread_messages( **message.model_dump(), 'reply_count': 0, 'latest_reply_at': None, - 'reactions': Messages.get_reactions_by_message_id(message.id, db=db), + 'reactions': await Messages.get_reactions_by_message_id(message.id, db=db), 'user': user_info, } ) @@ -1331,14 +1331,14 @@ async def update_message_by_id( message_id: str, form_data: MessageForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) - channel = Channels.get_channel_by_id(id, db=db) + await check_channels_access(request, user) + channel = await Channels.get_channel_by_id(id, db=db) if not channel: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) - message = Messages.get_message_by_id(message_id, db=db) + message = await Messages.get_message_by_id(message_id, db=db) if not message: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) @@ -1346,19 +1346,19 @@ async def update_message_by_id( raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.DEFAULT()) if channel.type in ['group', 'dm']: - if not Channels.is_user_channel_member(channel.id, user.id, db=db): + if not await Channels.is_user_channel_member(channel.id, user.id, db=db): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) else: if ( user.role != 'admin' and message.user_id != user.id - and not channel_has_access(user.id, channel, permission='write', strict=False, db=db) + and not await channel_has_access(user.id, channel, permission='write', strict=False, db=db) ): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) try: - message = Messages.update_message_by_id(message_id, form_data, db=db) - message = Messages.get_message_by_id(message_id, db=db) + await Messages.update_message_by_id(message_id, form_data, db=db) + message = await Messages.get_message_by_id(message_id, db=db) if message: await sio.emit( @@ -1398,18 +1398,18 @@ async def add_reaction_to_message( message_id: str, form_data: ReactionForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) - channel = Channels.get_channel_by_id(id, db=db) + await check_channels_access(request, user) + channel = await Channels.get_channel_by_id(id, db=db) if not channel: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) if channel.type in ['group', 'dm']: - if not Channels.is_user_channel_member(channel.id, user.id, db=db): + if not await Channels.is_user_channel_member(channel.id, user.id, db=db): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) else: - if user.role != 'admin' and not channel_has_access( + if user.role != 'admin' and not await channel_has_access( user.id, channel, permission='write', @@ -1418,7 +1418,7 @@ async def add_reaction_to_message( ): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) - message = Messages.get_message_by_id(message_id, db=db) + message = await Messages.get_message_by_id(message_id, db=db) if not message: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) @@ -1426,8 +1426,8 @@ async def add_reaction_to_message( raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.DEFAULT()) try: - Messages.add_reaction_to_message(message_id, user.id, form_data.name, db=db) - message = Messages.get_message_by_id(message_id, db=db) + await Messages.add_reaction_to_message(message_id, user.id, form_data.name, db=db) + message = await Messages.get_message_by_id(message_id, db=db) await sio.emit( 'events:channel', @@ -1465,18 +1465,18 @@ async def remove_reaction_by_id_and_user_id_and_name( message_id: str, form_data: ReactionForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) - channel = Channels.get_channel_by_id(id, db=db) + await check_channels_access(request, user) + channel = await Channels.get_channel_by_id(id, db=db) if not channel: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) if channel.type in ['group', 'dm']: - if not Channels.is_user_channel_member(channel.id, user.id, db=db): + if not await Channels.is_user_channel_member(channel.id, user.id, db=db): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) else: - if user.role != 'admin' and not channel_has_access( + if user.role != 'admin' and not await channel_has_access( user.id, channel, permission='write', @@ -1485,7 +1485,7 @@ async def remove_reaction_by_id_and_user_id_and_name( ): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) - message = Messages.get_message_by_id(message_id, db=db) + message = await Messages.get_message_by_id(message_id, db=db) if not message: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) @@ -1493,9 +1493,9 @@ async def remove_reaction_by_id_and_user_id_and_name( raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.DEFAULT()) try: - Messages.remove_reaction_by_id_and_user_id_and_name(message_id, user.id, form_data.name, db=db) + await Messages.remove_reaction_by_id_and_user_id_and_name(message_id, user.id, form_data.name, db=db) - message = Messages.get_message_by_id(message_id, db=db) + message = await Messages.get_message_by_id(message_id, db=db) await sio.emit( 'events:channel', @@ -1532,14 +1532,14 @@ async def delete_message_by_id( id: str, message_id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) - channel = Channels.get_channel_by_id(id, db=db) + await check_channels_access(request, user) + channel = await Channels.get_channel_by_id(id, db=db) if not channel: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) - message = Messages.get_message_by_id(message_id, db=db) + message = await Messages.get_message_by_id(message_id, db=db) if not message: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) @@ -1547,13 +1547,13 @@ async def delete_message_by_id( raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.DEFAULT()) if channel.type in ['group', 'dm']: - if not Channels.is_user_channel_member(channel.id, user.id, db=db): + if not await Channels.is_user_channel_member(channel.id, user.id, db=db): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) else: if ( user.role != 'admin' and message.user_id != user.id - and not channel_has_access( + and not await channel_has_access( user.id, channel, permission='write', @@ -1564,7 +1564,7 @@ async def delete_message_by_id( raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) try: - Messages.delete_message_by_id(message_id, db=db) + await Messages.delete_message_by_id(message_id, db=db) await sio.emit( 'events:channel', { @@ -1585,7 +1585,7 @@ async def delete_message_by_id( if message.parent_id: # If this message is a reply, emit to the parent message as well - parent_message = Messages.get_message_by_id(message.parent_id, db=db) + parent_message = await Messages.get_message_by_id(message.parent_id, db=db) if parent_message: await sio.emit( @@ -1615,9 +1615,9 @@ async def delete_message_by_id( @router.get('/webhooks/{webhook_id}/profile/image') -def get_webhook_profile_image(webhook_id: str, user=Depends(get_verified_user)): +async def get_webhook_profile_image(webhook_id: str, user=Depends(get_verified_user)): """Get webhook profile image by webhook ID.""" - webhook = Channels.get_webhook_by_id(webhook_id) + webhook = await Channels.get_webhook_by_id(webhook_id) if not webhook: # Return default favicon if webhook not found return FileResponse(f'{STATIC_DIR}/favicon.png') @@ -1653,18 +1653,18 @@ async def get_channel_webhooks( request: Request, id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) - channel = Channels.get_channel_by_id(id, db=db) + await check_channels_access(request, user) + channel = await Channels.get_channel_by_id(id, db=db) if not channel: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) # Only channel managers can view webhooks - if not Channels.is_user_channel_manager(channel.id, user.id, db=db) and user.role != 'admin': + if not await Channels.is_user_channel_manager(channel.id, user.id, db=db) and user.role != 'admin': raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.UNAUTHORIZED) - return Channels.get_webhooks_by_channel_id(id, db=db) + return await Channels.get_webhooks_by_channel_id(id, db=db) @router.post('/{id}/webhooks/create', response_model=ChannelWebhookModel) @@ -1673,18 +1673,18 @@ async def create_channel_webhook( id: str, form_data: ChannelWebhookForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) - channel = Channels.get_channel_by_id(id, db=db) + await check_channels_access(request, user) + channel = await Channels.get_channel_by_id(id, db=db) if not channel: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) # Only channel managers can create webhooks - if not Channels.is_user_channel_manager(channel.id, user.id, db=db) and user.role != 'admin': + if not await Channels.is_user_channel_manager(channel.id, user.id, db=db) and user.role != 'admin': raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.UNAUTHORIZED) - webhook = Channels.insert_webhook(id, user.id, form_data, db=db) + webhook = await Channels.insert_webhook(id, user.id, form_data, db=db) if not webhook: raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.DEFAULT()) @@ -1698,22 +1698,22 @@ async def update_channel_webhook( webhook_id: str, form_data: ChannelWebhookForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) - channel = Channels.get_channel_by_id(id, db=db) + await check_channels_access(request, user) + channel = await Channels.get_channel_by_id(id, db=db) if not channel: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) # Only channel managers can update webhooks - if not Channels.is_user_channel_manager(channel.id, user.id, db=db) and user.role != 'admin': + if not await Channels.is_user_channel_manager(channel.id, user.id, db=db) and user.role != 'admin': raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.UNAUTHORIZED) - webhook = Channels.get_webhook_by_id(webhook_id, db=db) + webhook = await Channels.get_webhook_by_id(webhook_id, db=db) if not webhook or webhook.channel_id != id: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) - updated = Channels.update_webhook_by_id(webhook_id, form_data, db=db) + updated = await Channels.update_webhook_by_id(webhook_id, form_data, db=db) if not updated: raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.DEFAULT()) @@ -1726,22 +1726,22 @@ async def delete_channel_webhook( id: str, webhook_id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - check_channels_access(request, user) - channel = Channels.get_channel_by_id(id, db=db) + await check_channels_access(request, user) + channel = await Channels.get_channel_by_id(id, db=db) if not channel: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) # Only channel managers can delete webhooks - if not Channels.is_user_channel_manager(channel.id, user.id, db=db) and user.role != 'admin': + if not await Channels.is_user_channel_manager(channel.id, user.id, db=db) and user.role != 'admin': raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.UNAUTHORIZED) - webhook = Channels.get_webhook_by_id(webhook_id, db=db) + webhook = await Channels.get_webhook_by_id(webhook_id, db=db) if not webhook or webhook.channel_id != id: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) - return Channels.delete_webhook_by_id(webhook_id, db=db) + return await Channels.delete_webhook_by_id(webhook_id, db=db) ############################ @@ -1759,25 +1759,25 @@ async def post_webhook_message( webhook_id: str, token: str, form_data: WebhookMessageForm, - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Public endpoint to post messages via webhook. No authentication required.""" - check_channels_access(request) + await check_channels_access(request) # Validate webhook - webhook = Channels.get_webhook_by_id_and_token(webhook_id, token, db=db) + webhook = await Channels.get_webhook_by_id_and_token(webhook_id, token, db=db) if not webhook: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail='Invalid webhook URL', ) - channel = Channels.get_channel_by_id(webhook.channel_id, db=db) + channel = await Channels.get_channel_by_id(webhook.channel_id, db=db) if not channel: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) # Create message with webhook identity stored in meta - message = Messages.insert_new_message( + message = await Messages.insert_new_message( MessageForm(content=form_data.content, meta={'webhook': {'id': webhook.id}}), webhook.channel_id, webhook.user_id, # Required for DB but webhook info in meta takes precedence @@ -1791,10 +1791,10 @@ async def post_webhook_message( ) # Update last_used_at - Channels.update_webhook_last_used_at(webhook_id, db=db) + await Channels.update_webhook_last_used_at(webhook_id, db=db) # Get full message and emit event - message = Messages.get_message_by_id(message.id, db=db) + message = await Messages.get_message_by_id(message.id, db=db) event_data = { 'channel_id': channel.id, diff --git a/backend/open_webui/routers/chats.py b/backend/open_webui/routers/chats.py index 2d12e02523..ba07937ed1 100644 --- a/backend/open_webui/routers/chats.py +++ b/backend/open_webui/routers/chats.py @@ -1,7 +1,7 @@ import json import logging from typing import Optional -from sqlalchemy.orm import Session +from sqlalchemy.ext.asyncio import AsyncSession import asyncio from fastapi.responses import StreamingResponse @@ -25,7 +25,7 @@ from open_webui.models.chats import ( ) from open_webui.models.tags import TagModel, Tags from open_webui.models.folders import Folders -from open_webui.internal.db import get_session +from open_webui.internal.db import get_async_session from open_webui.config import ENABLE_ADMIN_CHAT_ACCESS, ENABLE_ADMIN_EXPORT from open_webui.constants import ERROR_MESSAGES @@ -49,19 +49,19 @@ router = APIRouter() @router.get('/', response_model=list[ChatTitleIdResponse]) @router.get('/list', response_model=list[ChatTitleIdResponse]) -def get_session_user_chat_list( +async def get_session_user_chat_list( user=Depends(get_verified_user), page: Optional[int] = None, include_pinned: Optional[bool] = False, include_folders: Optional[bool] = False, - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): try: if page is not None: limit = 60 skip = (page - 1) * limit - return Chats.get_chat_title_id_list_by_user_id( + return await Chats.get_chat_title_id_list_by_user_id( user.id, include_folders=include_folders, include_pinned=include_pinned, @@ -70,7 +70,7 @@ def get_session_user_chat_list( db=db, ) else: - return Chats.get_chat_title_id_list_by_user_id( + return await Chats.get_chat_title_id_list_by_user_id( user.id, include_folders=include_folders, include_pinned=include_pinned, @@ -88,17 +88,17 @@ def get_session_user_chat_list( @router.get('/stats/usage', response_model=ChatUsageStatsListResponse) -def get_session_user_chat_usage_stats( +async def get_session_user_chat_usage_stats( items_per_page: Optional[int] = 50, page: Optional[int] = 1, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): try: limit = items_per_page skip = (page - 1) * limit - result = Chats.get_chats_by_user_id(user.id, skip=skip, limit=limit, db=db) + result = await Chats.get_chats_by_user_id(user.id, skip=skip, limit=limit, db=db) chats = result.items total = result.total @@ -332,11 +332,11 @@ def _process_chat_for_export(chat) -> Optional[ChatStatsExport]: return None -def calculate_chat_stats(user_id, skip=0, limit=10, filter=None): +async def calculate_chat_stats(user_id, skip=0, limit=10, filter=None): if filter is None: filter = {} - result = Chats.get_chats_by_user_id( + result = await Chats.get_chats_by_user_id( user_id, skip=skip, limit=limit, @@ -352,12 +352,12 @@ def calculate_chat_stats(user_id, skip=0, limit=10, filter=None): return chat_stats_export_list, result.total -def generate_chat_stats_jsonl_generator(user_id, filter): +async def generate_chat_stats_jsonl_generator(user_id, filter): """ - Synchronous generator for streaming chat stats export. + Async generator for streaming chat stats export. NOTE: We intentionally do NOT pass a shared db session here. Instead, we let - each batch create its own short-lived session via get_db_context(None). + each batch create its own short-lived session via get_async_db_context(None). This is critical for SQLite in low-resource environments because: 1. SQLite uses file-level locking 2. Holding a session open for the entire streaming duration blocks other requests @@ -368,12 +368,12 @@ def generate_chat_stats_jsonl_generator(user_id, filter): while True: # Each batch gets its own session that closes after the query - result = Chats.get_chats_by_user_id( + result = await Chats.get_chats_by_user_id( user_id, filter=filter, skip=skip, limit=limit, - db=None, # Let get_db_context create a fresh session per batch + db=None, # Let get_async_db_context create a fresh session per batch ) if not result.items: break @@ -421,7 +421,7 @@ async def export_chat_stats( limit = CHAT_EXPORT_PAGE_ITEM_COUNT skip = (page - 1) * limit - chat_stats_export_list, total = await asyncio.to_thread(calculate_chat_stats, user.id, skip, limit, filter) + chat_stats_export_list, total = await calculate_chat_stats(user.id, skip, limit, filter) return ChatStatsExportList(items=chat_stats_export_list, total=total, page=page) @@ -440,7 +440,7 @@ async def export_single_chat_stats( request: Request, chat_id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """ Export stats for exactly one chat by ID. @@ -454,7 +454,7 @@ async def export_single_chat_stats( ) try: - chat = Chats.get_chat_by_id(chat_id, db=db) + chat = await Chats.get_chat_by_id(chat_id, db=db) if not chat: raise HTTPException( @@ -469,8 +469,8 @@ async def export_single_chat_stats( detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - # Process the chat for export - chat_stats = await asyncio.to_thread(_process_chat_for_export, chat) + # Process the chat for export (pure computation, no DB) + chat_stats = _process_chat_for_export(chat) if not chat_stats: raise HTTPException( @@ -491,15 +491,15 @@ async def export_single_chat_stats( async def delete_all_user_chats( request: Request, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - if user.role == 'user' and not has_permission(user.id, 'chat.delete', request.app.state.config.USER_PERMISSIONS): + if user.role == 'user' and not await has_permission(user.id, 'chat.delete', request.app.state.config.USER_PERMISSIONS): raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - result = Chats.delete_chats_by_user_id(user.id, db=db) + result = await Chats.delete_chats_by_user_id(user.id, db=db) return result @@ -516,7 +516,7 @@ async def get_user_chat_list_by_user_id( order_by: Optional[str] = None, direction: Optional[str] = None, user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): if not ENABLE_ADMIN_CHAT_ACCESS: raise HTTPException( @@ -538,7 +538,7 @@ async def get_user_chat_list_by_user_id( if direction: filter['direction'] = direction - return Chats.get_chat_list_by_user_id(user_id, include_archived=True, filter=filter, skip=skip, limit=limit, db=db) + return await Chats.get_chat_list_by_user_id(user_id, include_archived=True, filter=filter, skip=skip, limit=limit, db=db) ############################ @@ -550,10 +550,10 @@ async def get_user_chat_list_by_user_id( async def create_new_chat( form_data: ChatForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): try: - chat = Chats.insert_new_chat(user.id, form_data, db=db) + chat = await Chats.insert_new_chat(user.id, form_data, db=db) return ChatResponse(**chat.model_dump()) except Exception as e: log.exception(e) @@ -569,10 +569,10 @@ async def create_new_chat( async def import_chats( form_data: ChatsImportForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): try: - chats = Chats.import_chats(user.id, form_data.chats, db=db) + chats = await Chats.import_chats(user.id, form_data.chats, db=db) return chats except Exception as e: log.exception(e) @@ -585,11 +585,11 @@ async def import_chats( @router.get('/search', response_model=list[ChatTitleIdResponse]) -def search_user_chats( +async def search_user_chats( text: str, page: Optional[int] = None, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): if page is None: page = 1 @@ -599,7 +599,7 @@ def search_user_chats( chat_list = [ ChatTitleIdResponse(**chat.model_dump()) - for chat in Chats.get_chats_by_user_id_and_search_text(user.id, text, skip=skip, limit=limit, db=db) + for chat in await Chats.get_chats_by_user_id_and_search_text(user.id, text, skip=skip, limit=limit, db=db) ] # Delete tag if no chat is found @@ -607,9 +607,9 @@ def search_user_chats( if page == 1 and len(words) == 1 and words[0].startswith('tag:'): tag_id = words[0].replace('tag:', '') if len(chat_list) == 0: - if Tags.get_tag_by_name_and_user_id(tag_id, user.id, db=db): + if await Tags.get_tag_by_name_and_user_id(tag_id, user.id, db=db): log.debug(f'deleting tag: {tag_id}') - Tags.delete_tag_by_name_and_user_id(tag_id, user.id, db=db) + await Tags.delete_tag_by_name_and_user_id(tag_id, user.id, db=db) return chat_list @@ -620,15 +620,15 @@ def search_user_chats( @router.get('/folder/{folder_id}', response_model=list[ChatResponse]) -async def get_chats_by_folder_id(folder_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): +async def get_chats_by_folder_id(folder_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): folder_ids = [folder_id] - children_folders = Folders.get_children_folders_by_id_and_user_id(folder_id, user.id, db=db) + children_folders = await Folders.get_children_folders_by_id_and_user_id(folder_id, user.id, db=db) if children_folders: folder_ids.extend([folder.id for folder in children_folders]) return [ ChatResponse(**chat.model_dump()) - for chat in Chats.get_chats_by_folder_ids_and_user_id(folder_ids, user.id, db=db) + for chat in await Chats.get_chats_by_folder_ids_and_user_id(folder_ids, user.id, db=db) ] @@ -637,13 +637,13 @@ async def get_chat_list_by_folder_id( folder_id: str, page: Optional[int] = 1, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): try: limit = 10 skip = (page - 1) * limit - chats = Chats.get_chats_by_folder_id_and_user_id(folder_id, user.id, skip=skip, limit=limit, db=db) + chats = await Chats.get_chats_by_folder_id_and_user_id(folder_id, user.id, skip=skip, limit=limit, db=db) return [ {'title': chat.title, 'id': chat.id, 'updated_at': chat.updated_at, 'last_read_at': chat.last_read_at} for chat in chats @@ -660,8 +660,8 @@ async def get_chat_list_by_folder_id( @router.get('/pinned', response_model=list[ChatTitleIdResponse]) -async def get_user_pinned_chats(user=Depends(get_verified_user), db: Session = Depends(get_session)): - return Chats.get_pinned_chats_by_user_id(user.id, db=db) +async def get_user_pinned_chats(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + return await Chats.get_pinned_chats_by_user_id(user.id, db=db) ############################ @@ -670,8 +670,8 @@ async def get_user_pinned_chats(user=Depends(get_verified_user), db: Session = D @router.get('/all', response_model=list[ChatResponse]) -async def get_user_chats(user=Depends(get_verified_user), db: Session = Depends(get_session)): - result = Chats.get_chats_by_user_id(user.id, db=db) +async def get_user_chats(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + result = await Chats.get_chats_by_user_id(user.id, db=db) return [ChatResponse(**chat.model_dump()) for chat in result.items] @@ -681,8 +681,8 @@ async def get_user_chats(user=Depends(get_verified_user), db: Session = Depends( @router.get('/all/archived', response_model=list[ChatResponse]) -async def get_user_archived_chats(user=Depends(get_verified_user), db: Session = Depends(get_session)): - return [ChatResponse(**chat.model_dump()) for chat in Chats.get_archived_chats_by_user_id(user.id, db=db)] +async def get_user_archived_chats(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + return [ChatResponse(**chat.model_dump()) for chat in await Chats.get_archived_chats_by_user_id(user.id, db=db)] ############################ @@ -691,9 +691,9 @@ async def get_user_archived_chats(user=Depends(get_verified_user), db: Session = @router.get('/all/tags', response_model=list[TagModel]) -async def get_all_user_tags(user=Depends(get_verified_user), db: Session = Depends(get_session)): +async def get_all_user_tags(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): try: - tags = Tags.get_tags_by_user_id(user.id, db=db) + tags = await Tags.get_tags_by_user_id(user.id, db=db) return tags except Exception as e: log.exception(e) @@ -706,13 +706,13 @@ async def get_all_user_tags(user=Depends(get_verified_user), db: Session = Depen @router.get('/all/db', response_model=list[ChatResponse]) -async def get_all_user_chats_in_db(user=Depends(get_admin_user), db: Session = Depends(get_session)): +async def get_all_user_chats_in_db(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): if not ENABLE_ADMIN_EXPORT: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - return [ChatResponse(**chat.model_dump()) for chat in Chats.get_chats(db=db)] + return [ChatResponse(**chat.model_dump()) for chat in await Chats.get_chats(db=db)] ############################ @@ -727,7 +727,7 @@ async def get_archived_session_user_chat_list( order_by: Optional[str] = None, direction: Optional[str] = None, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): if page is None: page = 1 @@ -743,7 +743,7 @@ async def get_archived_session_user_chat_list( if direction: filter['direction'] = direction - return Chats.get_archived_chat_list_by_user_id( + return await Chats.get_archived_chat_list_by_user_id( user.id, filter=filter, skip=skip, @@ -758,8 +758,8 @@ async def get_archived_session_user_chat_list( @router.post('/archive/all', response_model=bool) -async def archive_all_chats(user=Depends(get_verified_user), db: Session = Depends(get_session)): - return Chats.archive_all_chats_by_user_id(user.id, db=db) +async def archive_all_chats(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + return await Chats.archive_all_chats_by_user_id(user.id, db=db) ############################ @@ -768,8 +768,8 @@ async def archive_all_chats(user=Depends(get_verified_user), db: Session = Depen @router.post('/unarchive/all', response_model=bool) -async def unarchive_all_chats(user=Depends(get_verified_user), db: Session = Depends(get_session)): - return Chats.unarchive_all_chats_by_user_id(user.id, db=db) +async def unarchive_all_chats(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + return await Chats.unarchive_all_chats_by_user_id(user.id, db=db) ############################ @@ -784,7 +784,7 @@ async def get_shared_session_user_chat_list( order_by: Optional[str] = None, direction: Optional[str] = None, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): if page is None: page = 1 @@ -800,7 +800,7 @@ async def get_shared_session_user_chat_list( if direction: filter['direction'] = direction - return Chats.get_shared_chat_list_by_user_id( + return await Chats.get_shared_chat_list_by_user_id( user.id, filter=filter, skip=skip, @@ -815,14 +815,14 @@ async def get_shared_session_user_chat_list( @router.get('/share/{share_id}', response_model=Optional[ChatResponse]) -async def get_shared_chat_by_id(share_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): +async def get_shared_chat_by_id(share_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): if user.role == 'pending': raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.NOT_FOUND) if user.role == 'user' or (user.role == 'admin' and not ENABLE_ADMIN_CHAT_ACCESS): - chat = Chats.get_chat_by_share_id(share_id, db=db) + chat = await Chats.get_chat_by_share_id(share_id, db=db) elif user.role == 'admin' and ENABLE_ADMIN_CHAT_ACCESS: - chat = Chats.get_chat_by_id(share_id, db=db) + chat = await Chats.get_chat_by_id(share_id, db=db) if chat: return ChatResponse(**chat.model_dump()) @@ -849,11 +849,11 @@ class TagFilterForm(TagForm): async def get_user_chat_list_by_tag_name( form_data: TagFilterForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - chats = Chats.get_chat_list_by_user_id_and_tag_name(user.id, form_data.name, form_data.skip, form_data.limit, db=db) + chats = await Chats.get_chat_list_by_user_id_and_tag_name(user.id, form_data.name, form_data.skip, form_data.limit, db=db) if len(chats) == 0: - Tags.delete_tag_by_name_and_user_id(form_data.name, user.id, db=db) + await Tags.delete_tag_by_name_and_user_id(form_data.name, user.id, db=db) return chats @@ -864,8 +864,8 @@ async def get_user_chat_list_by_tag_name( @router.get('/{id}', response_model=Optional[ChatResponse]) -async def get_chat_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db) +async def get_chat_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db) if chat: return ChatResponse(**chat.model_dump()) @@ -884,12 +884,12 @@ async def update_chat_by_id( id: str, form_data: ChatForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db) + chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db) if chat: updated_chat = {**chat.chat, **form_data.chat} - chat = Chats.update_chat_by_id(id, updated_chat, db=db) + chat = await Chats.update_chat_by_id(id, updated_chat, db=db) return ChatResponse(**chat.model_dump()) else: raise HTTPException( @@ -911,9 +911,9 @@ async def update_chat_message_by_id( message_id: str, form_data: MessageForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - chat = Chats.get_chat_by_id(id, db=db) + chat = await Chats.get_chat_by_id(id, db=db) if not chat: raise HTTPException( @@ -927,7 +927,7 @@ async def update_chat_message_by_id( detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - chat = Chats.upsert_message_to_chat_by_id_and_message_id( + chat = await Chats.upsert_message_to_chat_by_id_and_message_id( id, message_id, { @@ -935,7 +935,7 @@ async def update_chat_message_by_id( }, ) - event_emitter = get_event_emitter( + event_emitter = await get_event_emitter( { 'user_id': user.id, 'chat_id': id, @@ -973,9 +973,9 @@ async def send_chat_message_event_by_id( message_id: str, form_data: EventForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - chat = Chats.get_chat_by_id(id, db=db) + chat = await Chats.get_chat_by_id(id, db=db) if not chat: raise HTTPException( @@ -989,7 +989,7 @@ async def send_chat_message_event_by_id( detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - event_emitter = get_event_emitter( + event_emitter = await get_event_emitter( { 'user_id': user.id, 'chat_id': id, @@ -1017,36 +1017,36 @@ async def delete_chat_by_id( request: Request, id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): if user.role == 'admin': - chat = Chats.get_chat_by_id(id, db=db) + chat = await Chats.get_chat_by_id(id, db=db) if not chat: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND, ) - Chats.delete_orphan_tags_for_user(chat.meta.get('tags', []), user.id, threshold=1, db=db) + await Chats.delete_orphan_tags_for_user(chat.meta.get('tags', []), user.id, threshold=1, db=db) - result = Chats.delete_chat_by_id(id, db=db) + result = await Chats.delete_chat_by_id(id, db=db) return result else: - if not has_permission(user.id, 'chat.delete', request.app.state.config.USER_PERMISSIONS): + if not await has_permission(user.id, 'chat.delete', request.app.state.config.USER_PERMISSIONS): raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db) + chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db) if not chat: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND, ) - Chats.delete_orphan_tags_for_user(chat.meta.get('tags', []), user.id, threshold=1, db=db) + await Chats.delete_orphan_tags_for_user(chat.meta.get('tags', []), user.id, threshold=1, db=db) - result = Chats.delete_chat_by_id_and_user_id(id, user.id, db=db) + result = await Chats.delete_chat_by_id_and_user_id(id, user.id, db=db) return result @@ -1056,8 +1056,8 @@ async def delete_chat_by_id( @router.get('/{id}/pinned', response_model=Optional[bool]) -async def get_pinned_status_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db) +async def get_pinned_status_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db) if chat: return chat.pinned else: @@ -1070,10 +1070,10 @@ async def get_pinned_status_by_id(id: str, user=Depends(get_verified_user), db: @router.post('/{id}/pin', response_model=Optional[ChatResponse]) -async def pin_chat_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db) +async def pin_chat_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db) if chat: - chat = Chats.toggle_chat_pinned_by_id(id, db=db) + chat = await Chats.toggle_chat_pinned_by_id(id, db=db) return chat else: raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.DEFAULT()) @@ -1093,9 +1093,9 @@ async def clone_chat_by_id( form_data: CloneForm, id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db) + chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db) if chat: updated_chat = { **chat.chat, @@ -1104,7 +1104,7 @@ async def clone_chat_by_id( 'title': form_data.title if form_data.title else f'Clone of {chat.title}', } - chats = Chats.import_chats( + chats = await Chats.import_chats( user.id, [ ChatImportForm( @@ -1137,11 +1137,11 @@ async def clone_chat_by_id( @router.post('/{id}/clone/shared', response_model=Optional[ChatResponse]) -async def clone_shared_chat_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): +async def clone_shared_chat_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): if user.role == 'admin': - chat = Chats.get_chat_by_id(id, db=db) + chat = await Chats.get_chat_by_id(id, db=db) else: - chat = Chats.get_chat_by_share_id(id, db=db) + chat = await Chats.get_chat_by_share_id(id, db=db) if chat: updated_chat = { @@ -1151,7 +1151,7 @@ async def clone_shared_chat_by_id(id: str, user=Depends(get_verified_user), db: 'title': f'Clone of {chat.title}', } - chats = Chats.import_chats( + chats = await Chats.import_chats( user.id, [ ChatImportForm( @@ -1184,18 +1184,18 @@ async def clone_shared_chat_by_id(id: str, user=Depends(get_verified_user), db: @router.post('/{id}/archive', response_model=Optional[ChatResponse]) -async def archive_chat_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db) +async def archive_chat_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db) if chat: - chat = Chats.toggle_chat_archive_by_id(id, db=db) + chat = await Chats.toggle_chat_archive_by_id(id, db=db) tag_ids = chat.meta.get('tags', []) if chat.archived: # Archived chats are excluded from count — clean up orphans - Chats.delete_orphan_tags_for_user(tag_ids, user.id, db=db) + await Chats.delete_orphan_tags_for_user(tag_ids, user.id, db=db) else: # Unarchived — ensure tag rows exist - Tags.ensure_tags_exist(tag_ids, user.id, db=db) + await Tags.ensure_tags_exist(tag_ids, user.id, db=db) return ChatResponse(**chat.model_dump()) else: @@ -1212,24 +1212,24 @@ async def share_chat_by_id( request: Request, id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): if (user.role != 'admin') and ( - not has_permission(user.id, 'chat.share', request.app.state.config.USER_PERMISSIONS) + not await has_permission(user.id, 'chat.share', request.app.state.config.USER_PERMISSIONS) ): raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db) + chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db) if chat: if chat.share_id: - shared_chat = Chats.update_shared_chat_by_chat_id(chat.id, db=db) + shared_chat = await Chats.update_shared_chat_by_chat_id(chat.id, db=db) return ChatResponse(**shared_chat.model_dump()) - shared_chat = Chats.insert_shared_chat_by_chat_id(chat.id, db=db) + shared_chat = await Chats.insert_shared_chat_by_chat_id(chat.id, db=db) if not shared_chat: raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, @@ -1250,14 +1250,14 @@ async def share_chat_by_id( @router.delete('/{id}/share', response_model=Optional[bool]) -async def delete_shared_chat_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db) +async def delete_shared_chat_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db) if chat: if not chat.share_id: return False - result = Chats.delete_shared_chat_by_chat_id(id, db=db) - update_result = Chats.update_chat_share_id_by_id(id, None, db=db) + result = await Chats.delete_shared_chat_by_chat_id(id, db=db) + update_result = await Chats.update_chat_share_id_by_id(id, None, db=db) return result and update_result != None else: @@ -1281,11 +1281,11 @@ async def update_chat_folder_id_by_id( id: str, form_data: ChatFolderIdForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db) + chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db) if chat: - chat = Chats.update_chat_folder_id_by_id_and_user_id(id, user.id, form_data.folder_id, db=db) + chat = await Chats.update_chat_folder_id_by_id_and_user_id(id, user.id, form_data.folder_id, db=db) return ChatResponse(**chat.model_dump()) else: raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.DEFAULT()) @@ -1297,11 +1297,11 @@ async def update_chat_folder_id_by_id( @router.get('/{id}/tags', response_model=list[TagModel]) -async def get_chat_tags_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db) +async def get_chat_tags_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db) if chat: tags = chat.meta.get('tags', []) - return Tags.get_tags_by_ids_and_user_id(tags, user.id, db=db) + return await Tags.get_tags_by_ids_and_user_id(tags, user.id, db=db) else: raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.NOT_FOUND) @@ -1316,9 +1316,9 @@ async def add_tag_by_id_and_tag_name( id: str, form_data: TagForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db) + chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db) if chat: tags = chat.meta.get('tags', []) tag_id = form_data.name.replace(' ', '_').lower() @@ -1330,11 +1330,11 @@ async def add_tag_by_id_and_tag_name( ) if tag_id not in tags: - Chats.add_chat_tag_by_id_and_user_id_and_tag_name(id, user.id, form_data.name, db=db) + await Chats.add_chat_tag_by_id_and_user_id_and_tag_name(id, user.id, form_data.name, db=db) - chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db) + chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db) tags = chat.meta.get('tags', []) - return Tags.get_tags_by_ids_and_user_id(tags, user.id, db=db) + return await Tags.get_tags_by_ids_and_user_id(tags, user.id, db=db) else: raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.DEFAULT()) @@ -1349,18 +1349,18 @@ async def delete_tag_by_id_and_tag_name( id: str, form_data: TagForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db) + chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db) if chat: - Chats.delete_tag_by_id_and_user_id_and_tag_name(id, user.id, form_data.name, db=db) + await Chats.delete_tag_by_id_and_user_id_and_tag_name(id, user.id, form_data.name, db=db) - if Chats.count_chats_by_tag_name_and_user_id(form_data.name, user.id, db=db) == 0: - Tags.delete_tag_by_name_and_user_id(form_data.name, user.id, db=db) + if await Chats.count_chats_by_tag_name_and_user_id(form_data.name, user.id, db=db) == 0: + await Tags.delete_tag_by_name_and_user_id(form_data.name, user.id, db=db) - chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db) + chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db) tags = chat.meta.get('tags', []) - return Tags.get_tags_by_ids_and_user_id(tags, user.id, db=db) + return await Tags.get_tags_by_ids_and_user_id(tags, user.id, db=db) else: raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.NOT_FOUND) @@ -1371,12 +1371,12 @@ async def delete_tag_by_id_and_tag_name( @router.delete('/{id}/tags/all', response_model=Optional[bool]) -async def delete_all_tags_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db) +async def delete_all_tags_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db) if chat: old_tags = chat.meta.get('tags', []) - Chats.delete_all_tags_by_id_and_user_id(id, user.id, db=db) - Chats.delete_orphan_tags_for_user(old_tags, user.id, db=db) + await Chats.delete_all_tags_by_id_and_user_id(id, user.id, db=db) + await Chats.delete_orphan_tags_for_user(old_tags, user.id, db=db) return True else: diff --git a/backend/open_webui/routers/evaluations.py b/backend/open_webui/routers/evaluations.py index 9805f2ece2..6a847c22a5 100644 --- a/backend/open_webui/routers/evaluations.py +++ b/backend/open_webui/routers/evaluations.py @@ -20,8 +20,8 @@ from open_webui.models.feedbacks import ( from open_webui.constants import ERROR_MESSAGES from open_webui.utils.auth import get_admin_user, get_verified_user -from open_webui.internal.db import get_session -from sqlalchemy.orm import Session +from open_webui.internal.db import get_async_session +from sqlalchemy.ext.asyncio import AsyncSession log = logging.getLogger(__name__) @@ -208,10 +208,10 @@ class LeaderboardResponse(BaseModel): async def get_leaderboard( query: Optional[str] = None, user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Get model leaderboard with Elo ratings. Query filters by tag similarity.""" - feedbacks = Feedbacks.get_feedbacks_for_leaderboard(db=db) + feedbacks = await Feedbacks.get_feedbacks_for_leaderboard(db=db) similarities = None if query and query.strip(): @@ -244,10 +244,10 @@ async def get_model_history( model_id: str, days: int = 30, user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Get daily win/loss history for a specific model.""" - history = Feedbacks.get_model_evaluation_history(model_id=model_id, days=days, db=db) + history = await Feedbacks.get_model_evaluation_history(model_id=model_id, days=days, db=db) return ModelHistoryResponse(model_id=model_id, history=history) @@ -292,24 +292,24 @@ async def update_config( @router.get('/feedbacks/models', response_model=list[str]) -async def get_feedback_model_ids(user=Depends(get_admin_user), db: Session = Depends(get_session)): - return Feedbacks.get_distinct_model_ids(db=db) +async def get_feedback_model_ids(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): + return await Feedbacks.get_distinct_model_ids(db=db) @router.get('/feedbacks/all', response_model=list[FeedbackResponse]) -async def get_all_feedbacks(user=Depends(get_admin_user), db: Session = Depends(get_session)): - feedbacks = Feedbacks.get_all_feedbacks(db=db) +async def get_all_feedbacks(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): + feedbacks = await Feedbacks.get_all_feedbacks(db=db) return feedbacks @router.get('/feedbacks/all/ids', response_model=list[FeedbackIdResponse]) -async def get_all_feedback_ids(user=Depends(get_admin_user), db: Session = Depends(get_session)): - return Feedbacks.get_all_feedback_ids(db=db) +async def get_all_feedback_ids(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): + return await Feedbacks.get_all_feedback_ids(db=db) @router.delete('/feedbacks/all') -async def delete_all_feedbacks(user=Depends(get_admin_user), db: Session = Depends(get_session)): - success = Feedbacks.delete_all_feedbacks(db=db) +async def delete_all_feedbacks(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): + success = await Feedbacks.delete_all_feedbacks(db=db) return success @@ -317,23 +317,23 @@ async def delete_all_feedbacks(user=Depends(get_admin_user), db: Session = Depen async def export_all_feedbacks( model_id: Optional[str] = None, user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - feedbacks = Feedbacks.get_all_feedbacks(db=db) + feedbacks = await Feedbacks.get_all_feedbacks(db=db) if model_id: feedbacks = [f for f in feedbacks if f.data and f.data.get('model_id') == model_id] return feedbacks @router.get('/feedbacks/user', response_model=list[FeedbackUserResponse]) -async def get_feedbacks(user=Depends(get_verified_user), db: Session = Depends(get_session)): - feedbacks = Feedbacks.get_feedbacks_by_user_id(user.id, db=db) +async def get_feedbacks(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + feedbacks = await Feedbacks.get_feedbacks_by_user_id(user.id, db=db) return feedbacks @router.delete('/feedbacks', response_model=bool) -async def delete_feedbacks(user=Depends(get_verified_user), db: Session = Depends(get_session)): - success = Feedbacks.delete_feedbacks_by_user_id(user.id, db=db) +async def delete_feedbacks(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + success = await Feedbacks.delete_feedbacks_by_user_id(user.id, db=db) return success @@ -347,7 +347,7 @@ async def get_feedbacks( page: Optional[int] = 1, model_id: Optional[str] = None, user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): limit = PAGE_ITEM_COUNT @@ -362,7 +362,7 @@ async def get_feedbacks( if model_id: filter['model_id'] = model_id - result = Feedbacks.get_feedback_items(filter=filter, skip=skip, limit=limit, db=db) + result = await Feedbacks.get_feedback_items(filter=filter, skip=skip, limit=limit, db=db) return result @@ -371,9 +371,9 @@ async def create_feedback( request: Request, form_data: FeedbackForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - feedback = Feedbacks.insert_new_feedback(user_id=user.id, form_data=form_data, db=db) + feedback = await Feedbacks.insert_new_feedback(user_id=user.id, form_data=form_data, db=db) if not feedback: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -384,11 +384,11 @@ async def create_feedback( @router.get('/feedback/{id}', response_model=FeedbackModel) -async def get_feedback_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): +async def get_feedback_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): if user.role == 'admin': - feedback = Feedbacks.get_feedback_by_id(id=id, db=db) + feedback = await Feedbacks.get_feedback_by_id(id=id, db=db) else: - feedback = Feedbacks.get_feedback_by_id_and_user_id(id=id, user_id=user.id, db=db) + feedback = await Feedbacks.get_feedback_by_id_and_user_id(id=id, user_id=user.id, db=db) if not feedback: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) @@ -401,12 +401,12 @@ async def update_feedback_by_id( id: str, form_data: FeedbackForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): if user.role == 'admin': - feedback = Feedbacks.update_feedback_by_id(id=id, form_data=form_data, db=db) + feedback = await Feedbacks.update_feedback_by_id(id=id, form_data=form_data, db=db) else: - feedback = Feedbacks.update_feedback_by_id_and_user_id(id=id, user_id=user.id, form_data=form_data, db=db) + feedback = await Feedbacks.update_feedback_by_id_and_user_id(id=id, user_id=user.id, form_data=form_data, db=db) if not feedback: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) @@ -415,11 +415,11 @@ async def update_feedback_by_id( @router.delete('/feedback/{id}') -async def delete_feedback_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): +async def delete_feedback_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): if user.role == 'admin': - success = Feedbacks.delete_feedback_by_id(id=id, db=db) + success = await Feedbacks.delete_feedback_by_id(id=id, db=db) else: - success = Feedbacks.delete_feedback_by_id_and_user_id(id=id, user_id=user.id, db=db) + success = await Feedbacks.delete_feedback_by_id_and_user_id(id=id, user_id=user.id, db=db) if not success: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) diff --git a/backend/open_webui/routers/files.py b/backend/open_webui/routers/files.py index 84227c0eca..8f1ee13f7f 100644 --- a/backend/open_webui/routers/files.py +++ b/backend/open_webui/routers/files.py @@ -21,8 +21,8 @@ from fastapi import ( ) from fastapi.responses import FileResponse, StreamingResponse -from sqlalchemy.orm import Session -from open_webui.internal.db import get_session, SessionLocal +from sqlalchemy.ext.asyncio import AsyncSession +from open_webui.internal.db import get_async_session, SessionLocal from open_webui.constants import ERROR_MESSAGES from open_webui.retrieval.vector.factory import VECTOR_DB_CLIENT @@ -88,16 +88,16 @@ def _is_text_file(file_path: str, chunk_size: int = 8192) -> bool: return False -def process_uploaded_file( +async def process_uploaded_file( request, file, file_path, file_item, file_metadata, user, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ): - def _process_handler(db_session): + async def _process_handler(db_session): try: content_type = file.content_type @@ -141,7 +141,7 @@ def process_uploaded_file( except Exception as e: log.error(f'Error processing file: {file_item.id}') - Files.update_file_data_by_id( + await Files.update_file_data_by_id( file_item.id, { 'status': 'failed', @@ -158,7 +158,7 @@ def process_uploaded_file( @router.post('/', response_model=FileModelResponse) -def upload_file( +async def upload_file( request: Request, background_tasks: BackgroundTasks, file: UploadFile = File(...), @@ -166,9 +166,9 @@ def upload_file( process: bool = Query(True), process_in_background: bool = Query(True), user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - return upload_file_handler( + return await upload_file_handler( request, file=file, metadata=metadata, @@ -180,7 +180,7 @@ def upload_file( ) -def upload_file_handler( +async def upload_file_handler( request: Request, file: UploadFile = File(...), metadata: Optional[dict | str] = Form(None), @@ -188,7 +188,7 @@ def upload_file_handler( process_in_background: bool = Query(True), user=Depends(get_verified_user), background_tasks: Optional[BackgroundTasks] = None, - db: Optional[Session] = None, + db: Optional[AsyncSession] = None, ): log.info(f'file.content_type: {file.content_type} {process}') @@ -236,7 +236,7 @@ def upload_file_handler( }, ) - file_item = Files.insert_new_file( + file_item = await Files.insert_new_file( user.id, FileForm( **{ @@ -258,9 +258,9 @@ def upload_file_handler( ) if 'channel_id' in file_metadata: - channel = Channels.get_channel_by_id_and_user_id(file_metadata['channel_id'], user.id, db=db) + channel = await Channels.get_channel_by_id_and_user_id(file_metadata['channel_id'], user.id, db=db) if channel: - Channels.add_file_to_channel_by_id(channel.id, file_item.id, user.id, db=db) + await Channels.add_file_to_channel_by_id(channel.id, file_item.id, user.id, db=db) if process: if background_tasks and process_in_background: @@ -275,7 +275,7 @@ def upload_file_handler( ) return {'status': True, **file_item.model_dump()} else: - process_uploaded_file( + await process_uploaded_file( request, file, file_path, @@ -317,12 +317,12 @@ async def list_files( user=Depends(get_verified_user), page: int = Query(1, ge=1, description='Page number (1-indexed)'), content: bool = Query(True), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): skip = (page - 1) * PAGE_SIZE user_id = None if (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL) else user.id - result = Files.get_file_list(user_id=user_id, skip=skip, limit=PAGE_SIZE, db=db) + result = await Files.get_file_list(user_id=user_id, skip=skip, limit=PAGE_SIZE, db=db) if not content: for file in result.items: @@ -347,7 +347,7 @@ async def search_files( skip: int = Query(0, ge=0, description='Number of files to skip'), limit: int = Query(100, ge=1, le=1000, description='Maximum number of files to return'), user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """ Search for files by filename with support for wildcard patterns. @@ -357,7 +357,7 @@ async def search_files( user_id = None if (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL) else user.id # Use optimized database query with pagination - files = Files.search_files( + files = await Files.search_files( user_id=user_id, filename=filename, skip=skip, @@ -385,8 +385,8 @@ async def search_files( @router.delete('/all') -async def delete_all_files(user=Depends(get_admin_user), db: Session = Depends(get_session)): - result = Files.delete_all_files(db=db) +async def delete_all_files(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): + result = await Files.delete_all_files(db=db) if result: try: Storage.delete_all_files() @@ -412,8 +412,8 @@ async def delete_all_files(user=Depends(get_admin_user), db: Session = Depends(g @router.get('/{id}', response_model=Optional[FileModel]) -async def get_file_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - file = Files.get_file_by_id(id, db=db) +async def get_file_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + file = await Files.get_file_by_id(id, db=db) if not file: raise HTTPException( @@ -421,7 +421,7 @@ async def get_file_by_id(id: str, user=Depends(get_verified_user), db: Session = detail=ERROR_MESSAGES.NOT_FOUND, ) - if file.user_id == user.id or user.role == 'admin' or has_access_to_file(id, 'read', user, db=db): + if file.user_id == user.id or user.role == 'admin' or await has_access_to_file(id, 'read', user, db=db): return file else: raise HTTPException( @@ -435,9 +435,9 @@ async def get_file_process_status( id: str, stream: bool = Query(False), user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - file = Files.get_file_by_id(id, db=db) + file = await Files.get_file_by_id(id, db=db) if not file: raise HTTPException( @@ -445,7 +445,7 @@ async def get_file_process_status( detail=ERROR_MESSAGES.NOT_FOUND, ) - if file.user_id == user.id or user.role == 'admin' or has_access_to_file(id, 'read', user, db=db): + if file.user_id == user.id or user.role == 'admin' or await has_access_to_file(id, 'read', user, db=db): if stream: MAX_FILE_PROCESSING_DURATION = 3600 * 2 @@ -454,7 +454,7 @@ async def get_file_process_status( # Each poll creates its own short-lived session to avoid holding a # connection for hours. A WebSocket push would be more efficient. for _ in range(MAX_FILE_PROCESSING_DURATION): - file_item = Files.get_file_by_id(file_id) # Creates own session + file_item = await Files.get_file_by_id(file_id) # Creates own session if file_item: data = file_item.model_dump().get('data', {}) status = data.get('status') @@ -495,8 +495,8 @@ async def get_file_process_status( @router.get('/{id}/data/content') -async def get_file_data_content_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - file = Files.get_file_by_id(id, db=db) +async def get_file_data_content_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + file = await Files.get_file_by_id(id, db=db) if not file: raise HTTPException( @@ -504,7 +504,7 @@ async def get_file_data_content_by_id(id: str, user=Depends(get_verified_user), detail=ERROR_MESSAGES.NOT_FOUND, ) - if file.user_id == user.id or user.role == 'admin' or has_access_to_file(id, 'read', user, db=db): + if file.user_id == user.id or user.role == 'admin' or await has_access_to_file(id, 'read', user, db=db): return {'content': file.data.get('content', '')} else: raise HTTPException( @@ -523,14 +523,14 @@ class ContentForm(BaseModel): @router.post('/{id}/data/content/update') -def update_file_data_content_by_id( +async def update_file_data_content_by_id( request: Request, id: str, form_data: ContentForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - file = Files.get_file_by_id(id, db=db) + file = await Files.get_file_by_id(id, db=db) if not file: raise HTTPException( @@ -538,7 +538,7 @@ def update_file_data_content_by_id( detail=ERROR_MESSAGES.NOT_FOUND, ) - if file.user_id == user.id or user.role == 'admin' or has_access_to_file(id, 'write', user, db=db): + if file.user_id == user.id or user.role == 'admin' or await has_access_to_file(id, 'write', user, db=db): try: process_file( request, @@ -546,7 +546,7 @@ def update_file_data_content_by_id( user=user, db=db, ) - file = Files.get_file_by_id(id=id, db=db) + file = await Files.get_file_by_id(id=id, db=db) except Exception as e: log.exception(e) log.error(f'Error processing file: {file.id}') @@ -554,7 +554,7 @@ def update_file_data_content_by_id( # Propagate content change to all knowledge collections referencing # this file. Without this the old embeddings remain in the knowledge # collection and RAG returns both stale and current data (#20558). - knowledges = Knowledges.get_knowledges_by_file_id(id, db=db) + knowledges = await Knowledges.get_knowledges_by_file_id(id, db=db) for knowledge in knowledges: try: # Remove old embeddings for this file from the KB collection @@ -587,9 +587,9 @@ async def get_file_content_by_id( id: str, user=Depends(get_verified_user), attachment: bool = Query(False), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - file = Files.get_file_by_id(id, db=db) + file = await Files.get_file_by_id(id, db=db) if not file: raise HTTPException( @@ -597,7 +597,7 @@ async def get_file_content_by_id( detail=ERROR_MESSAGES.NOT_FOUND, ) - if file.user_id == user.id or user.role == 'admin' or has_access_to_file(id, 'read', user, db=db): + if file.user_id == user.id or user.role == 'admin' or await has_access_to_file(id, 'read', user, db=db): try: file_path = Storage.get_file(file.path) file_path = Path(file_path) @@ -646,8 +646,8 @@ async def get_file_content_by_id( @router.get('/{id}/content/html') -async def get_html_file_content_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - file = Files.get_file_by_id(id, db=db) +async def get_html_file_content_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + file = await Files.get_file_by_id(id, db=db) if not file: raise HTTPException( @@ -655,14 +655,14 @@ async def get_html_file_content_by_id(id: str, user=Depends(get_verified_user), detail=ERROR_MESSAGES.NOT_FOUND, ) - file_user = Users.get_user_by_id(file.user_id, db=db) + file_user = await Users.get_user_by_id(file.user_id, db=db) if not file_user or file_user.role != 'admin': raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND, ) - if file.user_id == user.id or user.role == 'admin' or has_access_to_file(id, 'read', user, db=db): + if file.user_id == user.id or user.role == 'admin' or await has_access_to_file(id, 'read', user, db=db): try: file_path = Storage.get_file(file.path) file_path = Path(file_path) @@ -693,8 +693,8 @@ async def get_html_file_content_by_id(id: str, user=Depends(get_verified_user), @router.get('/{id}/content/{file_name}') -async def get_file_content_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - file = Files.get_file_by_id(id, db=db) +async def get_file_content_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + file = await Files.get_file_by_id(id, db=db) if not file: raise HTTPException( @@ -702,7 +702,7 @@ async def get_file_content_by_id(id: str, user=Depends(get_verified_user), db: S detail=ERROR_MESSAGES.NOT_FOUND, ) - if file.user_id == user.id or user.role == 'admin' or has_access_to_file(id, 'read', user, db=db): + if file.user_id == user.id or user.role == 'admin' or await has_access_to_file(id, 'read', user, db=db): file_path = file.path # Handle Unicode filenames @@ -749,8 +749,8 @@ async def get_file_content_by_id(id: str, user=Depends(get_verified_user), db: S @router.delete('/{id}') -async def delete_file_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - file = Files.get_file_by_id(id, db=db) +async def delete_file_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + file = await Files.get_file_by_id(id, db=db) if not file: raise HTTPException( @@ -758,12 +758,12 @@ async def delete_file_by_id(id: str, user=Depends(get_verified_user), db: Sessio detail=ERROR_MESSAGES.NOT_FOUND, ) - if file.user_id == user.id or user.role == 'admin' or has_access_to_file(id, 'write', user, db=db): + if file.user_id == user.id or user.role == 'admin' or await has_access_to_file(id, 'write', user, db=db): # Clean up KB associations and embeddings before deleting - knowledges = Knowledges.get_knowledges_by_file_id(id, db=db) + knowledges = await Knowledges.get_knowledges_by_file_id(id, db=db) for knowledge in knowledges: # Remove KB-file relationship - Knowledges.remove_file_from_knowledge_by_id(knowledge.id, id, db=db) + await Knowledges.remove_file_from_knowledge_by_id(knowledge.id, id, db=db) # Clean KB embeddings (same logic as /knowledge/{id}/file/remove) try: VECTOR_DB_CLIENT.delete(collection_name=knowledge.id, filter={'file_id': id}) @@ -772,7 +772,7 @@ async def delete_file_by_id(id: str, user=Depends(get_verified_user), db: Sessio except Exception as e: log.debug(f'KB embedding cleanup for {knowledge.id}: {e}') - result = Files.delete_file_by_id(id, db=db) + result = await Files.delete_file_by_id(id, db=db) if result: try: Storage.delete_file(file.path) diff --git a/backend/open_webui/routers/folders.py b/backend/open_webui/routers/folders.py index 0bf5a87f1e..9938d1eca1 100644 --- a/backend/open_webui/routers/folders.py +++ b/backend/open_webui/routers/folders.py @@ -22,8 +22,8 @@ from open_webui.models.knowledge import Knowledges from open_webui.config import UPLOAD_DIR from open_webui.constants import ERROR_MESSAGES -from open_webui.internal.db import get_session -from sqlalchemy.orm import Session +from open_webui.internal.db import get_async_session +from sqlalchemy.ext.asyncio import AsyncSession from fastapi import APIRouter, Depends, File, HTTPException, UploadFile, status, Request @@ -48,7 +48,7 @@ router = APIRouter() async def get_folders( request: Request, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): if request.app.state.config.ENABLE_FOLDERS is False: raise HTTPException( @@ -56,7 +56,7 @@ async def get_folders( detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - if user.role != 'admin' and not has_permission( + if user.role != 'admin' and not await has_permission( user.id, 'features.folders', request.app.state.config.USER_PERMISSIONS, @@ -67,29 +67,29 @@ async def get_folders( detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - folders = Folders.get_folders_by_user_id(user.id, db=db) + folders = await Folders.get_folders_by_user_id(user.id, db=db) # Verify folder data integrity folder_list = [] for folder in folders: - if folder.parent_id and not Folders.get_folder_by_id_and_user_id(folder.parent_id, user.id, db=db): - folder = Folders.update_folder_parent_id_by_id_and_user_id(folder.id, user.id, None, db=db) + if folder.parent_id and not await Folders.get_folder_by_id_and_user_id(folder.parent_id, user.id, db=db): + folder = await Folders.update_folder_parent_id_by_id_and_user_id(folder.id, user.id, None, db=db) if folder.data: if 'files' in folder.data: valid_files = [] for file in folder.data['files']: if file.get('type') == 'file': - if Files.check_access_by_user_id(file.get('id'), user.id, 'read', db=db): + if await Files.check_access_by_user_id(file.get('id'), user.id, 'read', db=db): valid_files.append(file) elif file.get('type') == 'collection': - if Knowledges.check_access_by_user_id(file.get('id'), user.id, 'read', db=db): + if await Knowledges.check_access_by_user_id(file.get('id'), user.id, 'read', db=db): valid_files.append(file) else: valid_files.append(file) folder.data['files'] = valid_files - Folders.update_folder_by_id_and_user_id(folder.id, user.id, FolderUpdateForm(data=folder.data), db=db) + await Folders.update_folder_by_id_and_user_id(folder.id, user.id, FolderUpdateForm(data=folder.data), db=db) folder_list.append(FolderNameIdResponse(**folder.model_dump())) @@ -102,12 +102,12 @@ async def get_folders( @router.post('/') -def create_folder( +async def create_folder( form_data: FolderForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - folder = Folders.get_folder_by_parent_id_and_user_id_and_name(form_data.parent_id, user.id, form_data.name, db=db) + folder = await Folders.get_folder_by_parent_id_and_user_id_and_name(form_data.parent_id, user.id, form_data.name, db=db) if folder: raise HTTPException( @@ -116,7 +116,7 @@ def create_folder( ) try: - folder = Folders.insert_new_folder(user.id, form_data, form_data.parent_id, db=db) + folder = await Folders.insert_new_folder(user.id, form_data, form_data.parent_id, db=db) return folder except Exception as e: log.exception(e) @@ -133,8 +133,8 @@ def create_folder( @router.get('/{id}', response_model=Optional[FolderModel]) -async def get_folder_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - folder = Folders.get_folder_by_id_and_user_id(id, user.id, db=db) +async def get_folder_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + folder = await Folders.get_folder_by_id_and_user_id(id, user.id, db=db) if folder: return folder else: @@ -154,13 +154,13 @@ async def update_folder_name_by_id( id: str, form_data: FolderUpdateForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - folder = Folders.get_folder_by_id_and_user_id(id, user.id, db=db) + folder = await Folders.get_folder_by_id_and_user_id(id, user.id, db=db) if folder: if form_data.name is not None: # Check if folder with same name exists - existing_folder = Folders.get_folder_by_parent_id_and_user_id_and_name( + existing_folder = await Folders.get_folder_by_parent_id_and_user_id_and_name( folder.parent_id, user.id, form_data.name, db=db ) if existing_folder and existing_folder.id != id: @@ -170,7 +170,7 @@ async def update_folder_name_by_id( ) try: - folder = Folders.update_folder_by_id_and_user_id(id, user.id, form_data, db=db) + folder = await Folders.update_folder_by_id_and_user_id(id, user.id, form_data, db=db) return folder except Exception as e: log.exception(e) @@ -200,11 +200,11 @@ async def update_folder_parent_id_by_id( id: str, form_data: FolderParentIdForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - folder = Folders.get_folder_by_id_and_user_id(id, user.id, db=db) + folder = await Folders.get_folder_by_id_and_user_id(id, user.id, db=db) if folder: - existing_folder = Folders.get_folder_by_parent_id_and_user_id_and_name( + existing_folder = await Folders.get_folder_by_parent_id_and_user_id_and_name( form_data.parent_id, user.id, folder.name, db=db ) @@ -215,7 +215,7 @@ async def update_folder_parent_id_by_id( ) try: - folder = Folders.update_folder_parent_id_by_id_and_user_id(id, user.id, form_data.parent_id, db=db) + folder = await Folders.update_folder_parent_id_by_id_and_user_id(id, user.id, form_data.parent_id, db=db) return folder except Exception as e: log.exception(e) @@ -245,12 +245,12 @@ async def update_folder_is_expanded_by_id( id: str, form_data: FolderIsExpandedForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - folder = Folders.get_folder_by_id_and_user_id(id, user.id, db=db) + folder = await Folders.get_folder_by_id_and_user_id(id, user.id, db=db) if folder: try: - folder = Folders.update_folder_is_expanded_by_id_and_user_id(id, user.id, form_data.is_expanded, db=db) + folder = await Folders.update_folder_is_expanded_by_id_and_user_id(id, user.id, form_data.is_expanded, db=db) return folder except Exception as e: log.exception(e) @@ -277,10 +277,10 @@ async def delete_folder_by_id( id: str, delete_contents: Optional[bool] = True, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - if Chats.count_chats_by_folder_id_and_user_id(id, user.id, db=db): - chat_delete_permission = has_permission( + if await Chats.count_chats_by_folder_id_and_user_id(id, user.id, db=db): + chat_delete_permission = await has_permission( user.id, 'chat.delete', request.app.state.config.USER_PERMISSIONS, db=db ) if user.role != 'admin' and not chat_delete_permission: @@ -290,18 +290,18 @@ async def delete_folder_by_id( ) folders = [] - folders.append(Folders.get_folder_by_id_and_user_id(id, user.id, db=db)) + folders.append(await Folders.get_folder_by_id_and_user_id(id, user.id, db=db)) while folders: folder = folders.pop() if folder: try: - folder_ids = Folders.delete_folder_by_id_and_user_id(folder.id, user.id, db=db) + folder_ids = await Folders.delete_folder_by_id_and_user_id(folder.id, user.id, db=db) for folder_id in folder_ids: if delete_contents: - Chats.delete_chats_by_user_id_and_folder_id(user.id, folder_id, db=db) + await Chats.delete_chats_by_user_id_and_folder_id(user.id, folder_id, db=db) else: - Chats.move_chats_by_user_id_and_folder_id(user.id, folder_id, None, db=db) + await Chats.move_chats_by_user_id_and_folder_id(user.id, folder_id, None, db=db) return True except Exception as e: @@ -313,7 +313,7 @@ async def delete_folder_by_id( ) finally: # Get all subfolders - subfolders = Folders.get_folders_by_parent_id_and_user_id(folder.id, user.id, db=db) + subfolders = await Folders.get_folders_by_parent_id_and_user_id(folder.id, user.id, db=db) folders.extend(subfolders) else: diff --git a/backend/open_webui/routers/functions.py b/backend/open_webui/routers/functions.py index 01bcbc411c..371079bed2 100644 --- a/backend/open_webui/routers/functions.py +++ b/backend/open_webui/routers/functions.py @@ -26,8 +26,8 @@ from open_webui.constants import ERROR_MESSAGES from fastapi import APIRouter, Depends, HTTPException, Request, status from open_webui.utils.auth import get_admin_user, get_verified_user from pydantic import BaseModel, HttpUrl -from open_webui.internal.db import get_session -from sqlalchemy.orm import Session +from open_webui.internal.db import get_async_session +from sqlalchemy.ext.asyncio import AsyncSession log = logging.getLogger(__name__) @@ -42,13 +42,13 @@ router = APIRouter() @router.get('/', response_model=list[FunctionResponse]) -async def get_functions(user=Depends(get_verified_user), db: Session = Depends(get_session)): - return Functions.get_functions(db=db) +async def get_functions(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + return await Functions.get_functions(db=db) @router.get('/list', response_model=list[FunctionUserResponse]) -async def get_function_list(user=Depends(get_admin_user), db: Session = Depends(get_session)): - return Functions.get_function_list(db=db) +async def get_function_list(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): + return await Functions.get_function_list(db=db) ############################ @@ -60,9 +60,9 @@ async def get_function_list(user=Depends(get_admin_user), db: Session = Depends( async def get_functions( include_valves: bool = False, user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - return Functions.get_functions(include_valves=include_valves, db=db) + return await Functions.get_functions(include_valves=include_valves, db=db) ############################ @@ -145,12 +145,12 @@ async def sync_functions( request: Request, form_data: SyncFunctionsForm, user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): try: for function in form_data.functions: function.content = replace_imports(function.content) - function_module, function_type, frontmatter = load_function_module_by_id( + function_module, function_type, frontmatter = await load_function_module_by_id( function.id, content=function.content, ) @@ -163,7 +163,7 @@ async def sync_functions( log.exception(f'Error validating valves for function {function.id}: {e}') raise e - return Functions.sync_functions(user.id, form_data.functions, db=db) + return await Functions.sync_functions(user.id, form_data.functions, db=db) except Exception as e: log.exception(f'Failed to load a function: {e}') raise HTTPException( @@ -182,7 +182,7 @@ async def create_new_function( request: Request, form_data: FunctionForm, user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): if not form_data.id.isidentifier(): raise HTTPException( @@ -192,11 +192,11 @@ async def create_new_function( form_data.id = form_data.id.lower() - function = Functions.get_function_by_id(form_data.id, db=db) + function = await Functions.get_function_by_id(form_data.id, db=db) if function is None: try: form_data.content = replace_imports(form_data.content) - function_module, function_type, frontmatter = load_function_module_by_id( + function_module, function_type, frontmatter = await load_function_module_by_id( form_data.id, content=form_data.content, ) @@ -205,13 +205,13 @@ async def create_new_function( FUNCTIONS = request.app.state.FUNCTIONS FUNCTIONS[form_data.id] = function_module - function = Functions.insert_new_function(user.id, function_type, form_data, db=db) + function = await Functions.insert_new_function(user.id, function_type, form_data, db=db) function_cache_dir = CACHE_DIR / 'functions' / form_data.id function_cache_dir.mkdir(parents=True, exist_ok=True) if function_type == 'filter' and getattr(function_module, 'toggle', None): - Functions.update_function_metadata_by_id(form_data.id, {'toggle': True}, db=db) + await Functions.update_function_metadata_by_id(form_data.id, {'toggle': True}, db=db) if function: return function @@ -239,8 +239,8 @@ async def create_new_function( @router.get('/id/{id}', response_model=Optional[FunctionModel]) -async def get_function_by_id(id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)): - function = Functions.get_function_by_id(id, db=db) +async def get_function_by_id(id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): + function = await Functions.get_function_by_id(id, db=db) if function: return function @@ -257,10 +257,10 @@ async def get_function_by_id(id: str, user=Depends(get_admin_user), db: Session @router.post('/id/{id}/toggle', response_model=Optional[FunctionModel]) -async def toggle_function_by_id(id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)): - function = Functions.get_function_by_id(id, db=db) +async def toggle_function_by_id(id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): + function = await Functions.get_function_by_id(id, db=db) if function: - function = Functions.update_function_by_id(id, {'is_active': not function.is_active}, db=db) + function = await Functions.update_function_by_id(id, {'is_active': not function.is_active}, db=db) if function: return function @@ -282,10 +282,10 @@ async def toggle_function_by_id(id: str, user=Depends(get_admin_user), db: Sessi @router.post('/id/{id}/toggle/global', response_model=Optional[FunctionModel]) -async def toggle_global_by_id(id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)): - function = Functions.get_function_by_id(id, db=db) +async def toggle_global_by_id(id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): + function = await Functions.get_function_by_id(id, db=db) if function: - function = Functions.update_function_by_id(id, {'is_global': not function.is_global}, db=db) + function = await Functions.update_function_by_id(id, {'is_global': not function.is_global}, db=db) if function: return function @@ -312,11 +312,11 @@ async def update_function_by_id( id: str, form_data: FunctionForm, user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): try: form_data.content = replace_imports(form_data.content) - function_module, function_type, frontmatter = load_function_module_by_id(id, content=form_data.content) + function_module, function_type, frontmatter = await load_function_module_by_id(id, content=form_data.content) form_data.meta.manifest = frontmatter FUNCTIONS = request.app.state.FUNCTIONS @@ -325,10 +325,10 @@ async def update_function_by_id( updated = {**form_data.model_dump(exclude={'id'}), 'type': function_type} log.debug(updated) - function = Functions.update_function_by_id(id, updated, db=db) + function = await Functions.update_function_by_id(id, updated, db=db) if function_type == 'filter' and getattr(function_module, 'toggle', None): - Functions.update_function_metadata_by_id(id, {'toggle': True}, db=db) + await Functions.update_function_metadata_by_id(id, {'toggle': True}, db=db) if function: return function @@ -355,9 +355,9 @@ async def delete_function_by_id( request: Request, id: str, user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - result = Functions.delete_function_by_id(id, db=db) + result = await Functions.delete_function_by_id(id, db=db) if result: FUNCTIONS = request.app.state.FUNCTIONS @@ -373,11 +373,11 @@ async def delete_function_by_id( @router.get('/id/{id}/valves', response_model=Optional[dict]) -async def get_function_valves_by_id(id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)): - function = Functions.get_function_by_id(id, db=db) +async def get_function_valves_by_id(id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): + function = await Functions.get_function_by_id(id, db=db) if function: try: - valves = Functions.get_function_valves_by_id(id, db=db) + valves = await Functions.get_function_valves_by_id(id, db=db) return valves except Exception as e: raise HTTPException( @@ -401,11 +401,11 @@ async def get_function_valves_spec_by_id( request: Request, id: str, user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - function = Functions.get_function_by_id(id, db=db) + function = await Functions.get_function_by_id(id, db=db) if function: - function_module, function_type, frontmatter = get_function_module_from_cache(request, id) + function_module, function_type, frontmatter = await get_function_module_from_cache(request, id) if hasattr(function_module, 'Valves'): Valves = function_module.Valves @@ -432,11 +432,11 @@ async def update_function_valves_by_id( id: str, form_data: dict, user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - function = Functions.get_function_by_id(id, db=db) + function = await Functions.get_function_by_id(id, db=db) if function: - function_module, function_type, frontmatter = get_function_module_from_cache(request, id) + function_module, function_type, frontmatter = await get_function_module_from_cache(request, id) if hasattr(function_module, 'Valves'): Valves = function_module.Valves @@ -446,7 +446,7 @@ async def update_function_valves_by_id( valves = Valves(**form_data) valves_dict = valves.model_dump(exclude_unset=True) - Functions.update_function_valves_by_id(id, valves_dict, db=db) + await Functions.update_function_valves_by_id(id, valves_dict, db=db) return valves_dict except Exception as e: log.exception(f'Error updating function values by id {id}: {e}') @@ -473,11 +473,11 @@ async def update_function_valves_by_id( @router.get('/id/{id}/valves/user', response_model=Optional[dict]) -async def get_function_user_valves_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - function = Functions.get_function_by_id(id, db=db) +async def get_function_user_valves_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + function = await Functions.get_function_by_id(id, db=db) if function: try: - user_valves = Functions.get_user_valves_by_id_and_user_id(id, user.id, db=db) + user_valves = await Functions.get_user_valves_by_id_and_user_id(id, user.id, db=db) return user_valves except Exception as e: raise HTTPException( @@ -496,11 +496,11 @@ async def get_function_user_valves_spec_by_id( request: Request, id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - function = Functions.get_function_by_id(id, db=db) + function = await Functions.get_function_by_id(id, db=db) if function: - function_module, function_type, frontmatter = get_function_module_from_cache(request, id) + function_module, function_type, frontmatter = await get_function_module_from_cache(request, id) if hasattr(function_module, 'UserValves'): UserValves = function_module.UserValves @@ -522,12 +522,12 @@ async def update_function_user_valves_by_id( id: str, form_data: dict, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - function = Functions.get_function_by_id(id, db=db) + function = await Functions.get_function_by_id(id, db=db) if function: - function_module, function_type, frontmatter = get_function_module_from_cache(request, id) + function_module, function_type, frontmatter = await get_function_module_from_cache(request, id) if hasattr(function_module, 'UserValves'): UserValves = function_module.UserValves @@ -536,7 +536,7 @@ async def update_function_user_valves_by_id( form_data = {k: v for k, v in form_data.items() if v is not None} user_valves = UserValves(**form_data) user_valves_dict = user_valves.model_dump(exclude_unset=True) - Functions.update_user_valves_by_id_and_user_id(id, user.id, user_valves_dict, db=db) + await Functions.update_user_valves_by_id_and_user_id(id, user.id, user_valves_dict, db=db) return user_valves_dict except Exception as e: log.exception(f'Error updating function user valves by id {id}: {e}') diff --git a/backend/open_webui/routers/groups.py b/backend/open_webui/routers/groups.py index 4e9688c3d8..c45690fc3a 100755 --- a/backend/open_webui/routers/groups.py +++ b/backend/open_webui/routers/groups.py @@ -17,8 +17,8 @@ from open_webui.config import CACHE_DIR from open_webui.constants import ERROR_MESSAGES from fastapi import APIRouter, Depends, HTTPException, Request, status -from open_webui.internal.db import get_session -from sqlalchemy.orm import Session +from open_webui.internal.db import get_async_session +from sqlalchemy.ext.asyncio import AsyncSession from open_webui.utils.auth import get_admin_user, get_verified_user @@ -35,7 +35,7 @@ router = APIRouter() async def get_groups( share: Optional[bool] = None, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): filter = {} @@ -45,7 +45,7 @@ async def get_groups( if share is not None: filter['share'] = share - groups = Groups.get_groups(filter=filter, db=db) + groups = await Groups.get_groups(filter=filter, db=db) return groups @@ -59,14 +59,14 @@ async def get_groups( async def create_new_group( form_data: GroupForm, user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): try: - group = Groups.insert_new_group(user.id, form_data, db=db) + group = await Groups.insert_new_group(user.id, form_data, db=db) if group: return GroupResponse( **group.model_dump(), - member_count=Groups.get_group_member_count_by_id(group.id, db=db), + member_count=await Groups.get_group_member_count_by_id(group.id, db=db), ) else: raise HTTPException( @@ -87,12 +87,12 @@ async def create_new_group( @router.get('/id/{id}', response_model=Optional[GroupResponse]) -async def get_group_by_id(id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)): - group = Groups.get_group_by_id(id, db=db) +async def get_group_by_id(id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): + group = await Groups.get_group_by_id(id, db=db) if group: return GroupResponse( **group.model_dump(), - member_count=Groups.get_group_member_count_by_id(group.id, db=db), + member_count=await Groups.get_group_member_count_by_id(group.id, db=db), ) else: raise HTTPException( @@ -102,12 +102,12 @@ async def get_group_by_id(id: str, user=Depends(get_admin_user), db: Session = D @router.get('/id/{id}/info', response_model=Optional[GroupInfoResponse]) -async def get_group_info_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - group = Groups.get_group_by_id(id, db=db) +async def get_group_info_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + group = await Groups.get_group_by_id(id, db=db) if group: return GroupInfoResponse( **group.model_dump(), - member_count=Groups.get_group_member_count_by_id(group.id, db=db), + member_count=await Groups.get_group_member_count_by_id(group.id, db=db), ) else: raise HTTPException( @@ -127,13 +127,13 @@ class GroupExportResponse(GroupResponse): @router.get('/id/{id}/export', response_model=Optional[GroupExportResponse]) -async def export_group_by_id(id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)): - group = Groups.get_group_by_id(id, db=db) +async def export_group_by_id(id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): + group = await Groups.get_group_by_id(id, db=db) if group: return GroupExportResponse( **group.model_dump(), - member_count=Groups.get_group_member_count_by_id(group.id, db=db), - user_ids=Groups.get_group_user_ids_by_id(group.id, db=db), + member_count=await Groups.get_group_member_count_by_id(group.id, db=db), + user_ids=await Groups.get_group_user_ids_by_id(group.id, db=db), ) else: raise HTTPException( @@ -148,9 +148,9 @@ async def export_group_by_id(id: str, user=Depends(get_admin_user), db: Session @router.post('/id/{id}/users', response_model=list[UserInfoResponse]) -async def get_users_in_group(id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)): +async def get_users_in_group(id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): try: - users = Users.get_users_by_group_id(id, db=db) + users = await Users.get_users_by_group_id(id, db=db) return users except Exception as e: log.exception(f'Error adding users to group {id}: {e}') @@ -170,14 +170,14 @@ async def update_group_by_id( id: str, form_data: GroupUpdateForm, user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): try: - group = Groups.update_group_by_id(id, form_data, db=db) + group = await Groups.update_group_by_id(id, form_data, db=db) if group: return GroupResponse( **group.model_dump(), - member_count=Groups.get_group_member_count_by_id(group.id, db=db), + member_count=await Groups.get_group_member_count_by_id(group.id, db=db), ) else: raise HTTPException( @@ -202,17 +202,17 @@ async def add_user_to_group( id: str, form_data: UserIdsForm, user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): try: if form_data.user_ids: - form_data.user_ids = Users.get_valid_user_ids(form_data.user_ids, db=db) + form_data.user_ids = await Users.get_valid_user_ids(form_data.user_ids, db=db) - group = Groups.add_users_to_group(id, form_data.user_ids, db=db) + group = await Groups.add_users_to_group(id, form_data.user_ids, db=db) if group: return GroupResponse( **group.model_dump(), - member_count=Groups.get_group_member_count_by_id(group.id, db=db), + member_count=await Groups.get_group_member_count_by_id(group.id, db=db), ) else: raise HTTPException( @@ -232,14 +232,14 @@ async def remove_users_from_group( id: str, form_data: UserIdsForm, user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): try: - group = Groups.remove_users_from_group(id, form_data.user_ids, db=db) + group = await Groups.remove_users_from_group(id, form_data.user_ids, db=db) if group: return GroupResponse( **group.model_dump(), - member_count=Groups.get_group_member_count_by_id(group.id, db=db), + member_count=await Groups.get_group_member_count_by_id(group.id, db=db), ) else: raise HTTPException( @@ -260,9 +260,9 @@ async def remove_users_from_group( @router.delete('/id/{id}/delete', response_model=bool) -async def delete_group_by_id(id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)): +async def delete_group_by_id(id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): try: - result = Groups.delete_group_by_id(id, db=db) + result = await Groups.delete_group_by_id(id, db=db) if result: return result else: diff --git a/backend/open_webui/routers/images.py b/backend/open_webui/routers/images.py index 0e56da560b..fb3bc1cec5 100644 --- a/backend/open_webui/routers/images.py +++ b/backend/open_webui/routers/images.py @@ -28,8 +28,8 @@ from open_webui.routers.files import upload_file_handler, get_file_content_by_id from open_webui.utils.auth import get_admin_user, get_verified_user from open_webui.utils.access_control import has_permission from open_webui.utils.headers import include_user_info_headers -from open_webui.internal.db import get_session -from sqlalchemy.orm import Session +from open_webui.internal.db import get_async_session +from sqlalchemy.ext.asyncio import AsyncSession from open_webui.utils.images.comfyui import ( ComfyUICreateImageForm, ComfyUIEditImageForm, @@ -341,7 +341,7 @@ async def verify_url(request: Request, user=Depends(get_admin_user)): @router.get('/models') -def get_models(request: Request, user=Depends(get_verified_user)): +async def get_models(request: Request, user=Depends(get_verified_user)): try: if request.app.state.config.IMAGE_GENERATION_ENGINE == 'openai': return [ @@ -456,7 +456,7 @@ def get_image_data(data: str, headers=None): return None, None -def upload_image(request, image_data, content_type, metadata, user, db=None): +async def upload_image(request, image_data, content_type, metadata, user, db=None): image_format = mimetypes.guess_extension(content_type) file = UploadFile( file=io.BytesIO(image_data), @@ -465,7 +465,7 @@ def upload_image(request, image_data, content_type, metadata, user, db=None): 'content-type': content_type, }, ) - file_item = upload_file_handler( + file_item = await upload_file_handler( request, file=file, metadata=metadata, @@ -479,7 +479,7 @@ def upload_image(request, image_data, content_type, metadata, user, db=None): message_id = metadata.get('message_id') if chat_id and message_id: - Chats.insert_chat_files( + await Chats.insert_chat_files( chat_id=chat_id, message_id=message_id, file_ids=[file_item.id], @@ -499,7 +499,7 @@ async def generate_images(request: Request, form_data: CreateImageForm, user=Dep detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - if user.role != 'admin' and not has_permission( + if user.role != 'admin' and not await has_permission( user.id, 'features.image_generation', request.app.state.config.USER_PERMISSIONS ): raise HTTPException( @@ -590,7 +590,7 @@ async def image_generations( else: image_data, content_type = get_image_data(image['b64_json']) - _, url = upload_image(request, image_data, content_type, {**data, **metadata}, user) + _, url = await upload_image(request, image_data, content_type, {**data, **metadata}, user) images.append({'url': url}) return images @@ -635,14 +635,14 @@ async def image_generations( if model.endswith(':predict'): for image in res['predictions']: image_data, content_type = get_image_data(image['bytesBase64Encoded']) - _, url = upload_image(request, image_data, content_type, {**data, **metadata}, user) + _, url = await upload_image(request, image_data, content_type, {**data, **metadata}, user) images.append({'url': url}) elif model.endswith(':generateContent'): for image in res['candidates']: for part in image['content']['parts']: if part.get('inlineData', {}).get('data'): image_data, content_type = get_image_data(part['inlineData']['data']) - _, url = upload_image( + _, url = await upload_image( request, image_data, content_type, @@ -695,7 +695,7 @@ async def image_generations( headers = {'Authorization': f'Bearer {request.app.state.config.COMFYUI_API_KEY}'} image_data, content_type = get_image_data(image['url'], headers) - _, url = upload_image( + _, url = await upload_image( request, image_data, content_type, @@ -742,7 +742,7 @@ async def image_generations( for image in res['images']: image_data, content_type = get_image_data(image) - _, url = upload_image( + _, url = await upload_image( request, image_data, content_type, @@ -832,7 +832,7 @@ async def image_edits( except Exception as e: raise HTTPException(status_code=400, detail=ERROR_MESSAGES.DEFAULT(e)) - def get_image_file_item(base64_string, param_name='image'): + async def get_image_file_item(base64_string, param_name='image'): data = base64_string header, encoded = data.split(',', 1) mime_type = header.split(';')[0].lstrip('data:') @@ -905,7 +905,7 @@ async def image_edits( else: image_data, content_type = get_image_data(image['b64_json']) - _, url = upload_image(request, image_data, content_type, {**data, **metadata}, user) + _, url = await upload_image(request, image_data, content_type, {**data, **metadata}, user) images.append({'url': url}) return images @@ -956,7 +956,7 @@ async def image_edits( for part in image['content']['parts']: if part.get('inlineData', {}).get('data'): image_data, content_type = get_image_data(part['inlineData']['data']) - _, url = upload_image( + _, url = await upload_image( request, image_data, content_type, @@ -1036,7 +1036,7 @@ async def image_edits( headers = {'Authorization': f'Bearer {request.app.state.config.IMAGES_EDIT_COMFYUI_API_KEY}'} image_data, content_type = get_image_data(image_url, headers) - _, url = upload_image( + _, url = await upload_image( request, image_data, content_type, diff --git a/backend/open_webui/routers/knowledge.py b/backend/open_webui/routers/knowledge.py index ead782cdbf..f6c3416c8d 100644 --- a/backend/open_webui/routers/knowledge.py +++ b/backend/open_webui/routers/knowledge.py @@ -8,8 +8,8 @@ import io import zipfile from urllib.parse import quote -from sqlalchemy.orm import Session -from open_webui.internal.db import get_session +from sqlalchemy.ext.asyncio import AsyncSession +from open_webui.internal.db import get_async_session from open_webui.models.groups import Groups from open_webui.models.knowledge import ( KnowledgeFileListResponse, @@ -111,14 +111,14 @@ class KnowledgeAccessListResponse(BaseModel): async def get_knowledge_bases( page: Optional[int] = 1, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): page = max(page, 1) limit = PAGE_ITEM_COUNT skip = (page - 1) * limit filter = {} - groups = Groups.get_groups_by_member_id(user.id, db=db) + groups = await Groups.get_groups_by_member_id(user.id, db=db) user_group_ids = {group.id for group in groups} if not user.role == 'admin' or not BYPASS_ADMIN_ACCESS_CONTROL: @@ -127,11 +127,11 @@ async def get_knowledge_bases( filter['user_id'] = user.id - result = Knowledges.search_knowledge_bases(user.id, filter=filter, skip=skip, limit=limit, db=db) + result = await Knowledges.search_knowledge_bases(user.id, filter=filter, skip=skip, limit=limit, db=db) # Batch-fetch writable knowledge IDs in a single query instead of N has_access calls knowledge_base_ids = [knowledge_base.id for knowledge_base in result.items] - writable_knowledge_base_ids = AccessGrants.get_accessible_resource_ids( + writable_knowledge_base_ids = await AccessGrants.get_accessible_resource_ids( user_id=user.id, resource_type='knowledge', resource_ids=knowledge_base_ids, @@ -162,7 +162,7 @@ async def search_knowledge_bases( view_option: Optional[str] = None, page: Optional[int] = 1, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): page = max(page, 1) limit = PAGE_ITEM_COUNT @@ -174,7 +174,7 @@ async def search_knowledge_bases( if view_option: filter['view_option'] = view_option - groups = Groups.get_groups_by_member_id(user.id, db=db) + groups = await Groups.get_groups_by_member_id(user.id, db=db) user_group_ids = {group.id for group in groups} if not user.role == 'admin' or not BYPASS_ADMIN_ACCESS_CONTROL: @@ -183,11 +183,11 @@ async def search_knowledge_bases( filter['user_id'] = user.id - result = Knowledges.search_knowledge_bases(user.id, filter=filter, skip=skip, limit=limit, db=db) + result = await Knowledges.search_knowledge_bases(user.id, filter=filter, skip=skip, limit=limit, db=db) # Batch-fetch writable knowledge IDs in a single query instead of N has_access calls knowledge_base_ids = [knowledge_base.id for knowledge_base in result.items] - writable_knowledge_base_ids = AccessGrants.get_accessible_resource_ids( + writable_knowledge_base_ids = await AccessGrants.get_accessible_resource_ids( user_id=user.id, resource_type='knowledge', resource_ids=knowledge_base_ids, @@ -217,7 +217,7 @@ async def search_knowledge_files( query: Optional[str] = None, page: Optional[int] = 1, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): page = max(page, 1) limit = PAGE_ITEM_COUNT @@ -227,13 +227,13 @@ async def search_knowledge_files( if query: filter['query'] = query - groups = Groups.get_groups_by_member_id(user.id, db=db) + groups = await Groups.get_groups_by_member_id(user.id, db=db) if groups: filter['group_ids'] = [group.id for group in groups] filter['user_id'] = user.id - return Knowledges.search_knowledge_files(filter=filter, skip=skip, limit=limit, db=db) + return await Knowledges.search_knowledge_files(filter=filter, skip=skip, limit=limit, db=db) ############################ @@ -247,11 +247,11 @@ async def create_new_knowledge( form_data: KnowledgeForm, user=Depends(get_verified_user), ): - # NOTE: We intentionally do NOT use Depends(get_session) here. + # NOTE: We intentionally do NOT use Depends(get_async_session) here. # Database operations (has_permission, filter_allowed_access_grants, insert_new_knowledge) manage their own sessions. # This prevents holding a connection during embed_knowledge_base_metadata() # which makes external embedding API calls (1-5+ seconds). - if user.role != 'admin' and not has_permission( + if user.role != 'admin' and not await has_permission( user.id, 'workspace.knowledge', request.app.state.config.USER_PERMISSIONS ): raise HTTPException( @@ -259,7 +259,7 @@ async def create_new_knowledge( detail=ERROR_MESSAGES.UNAUTHORIZED, ) - form_data.access_grants = filter_allowed_access_grants( + form_data.access_grants = await filter_allowed_access_grants( request.app.state.config.USER_PERMISSIONS, user.id, user.role, @@ -267,7 +267,7 @@ async def create_new_knowledge( 'sharing.public_knowledge', ) - knowledge = Knowledges.insert_new_knowledge(user.id, form_data) + knowledge = await Knowledges.insert_new_knowledge(user.id, form_data) if knowledge: # Embed knowledge base for semantic search @@ -294,7 +294,7 @@ async def create_new_knowledge( async def reindex_knowledge_files( request: Request, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): if user.role != 'admin': raise HTTPException( @@ -302,13 +302,13 @@ async def reindex_knowledge_files( detail=ERROR_MESSAGES.UNAUTHORIZED, ) - knowledge_bases = Knowledges.get_knowledge_bases(db=db) + knowledge_bases = await Knowledges.get_knowledge_bases(db=db) log.info(f'Starting reindexing for {len(knowledge_bases)} knowledge bases') for knowledge_base in knowledge_bases: try: - files = Knowledges.get_files_by_id(knowledge_base.id, db=db) + files = await Knowledges.get_files_by_id(knowledge_base.id, db=db) try: if VECTOR_DB_CLIENT.has_collection(collection_name=knowledge_base.id): VECTOR_DB_CLIENT.delete_collection(collection_name=knowledge_base.id) @@ -357,12 +357,12 @@ async def reindex_knowledge_base_metadata_embeddings( ): """Batch embed all existing knowledge bases. Admin only. - NOTE: We intentionally do NOT use Depends(get_session) here. + NOTE: We intentionally do NOT use Depends(get_async_session) here. This endpoint loops through ALL knowledge bases and calls embed_knowledge_base_metadata() for each one, making N external embedding API calls. Holding a session during this entire operation would exhaust the connection pool. """ - knowledge_bases = Knowledges.get_knowledge_bases() + knowledge_bases = await Knowledges.get_knowledge_bases() log.info(f'Reindexing embeddings for {len(knowledge_bases)} knowledge bases') success_count = 0 @@ -385,14 +385,14 @@ class KnowledgeFilesResponse(KnowledgeResponse): @router.get('/{id}', response_model=Optional[KnowledgeFilesResponse]) -async def get_knowledge_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - knowledge = Knowledges.get_knowledge_by_id(id=id, db=db) +async def get_knowledge_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db) if knowledge: if ( user.role == 'admin' or knowledge.user_id == user.id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='knowledge', resource_id=knowledge.id, @@ -405,7 +405,7 @@ async def get_knowledge_by_id(id: str, user=Depends(get_verified_user), db: Sess write_access=( user.id == knowledge.user_id or (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL) - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='knowledge', resource_id=knowledge.id, @@ -438,11 +438,11 @@ async def update_knowledge_by_id( form_data: KnowledgeForm, user=Depends(get_verified_user), ): - # NOTE: We intentionally do NOT use Depends(get_session) here. + # NOTE: We intentionally do NOT use Depends(get_async_session) here. # Database operations manage their own short-lived sessions internally. # This prevents holding a connection during embed_knowledge_base_metadata() # which makes external embedding API calls (1-5+ seconds). - knowledge = Knowledges.get_knowledge_by_id(id=id) + knowledge = await Knowledges.get_knowledge_by_id(id=id) if not knowledge: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -451,7 +451,7 @@ async def update_knowledge_by_id( # Is the user the original creator, in a group with write access, or an admin if ( knowledge.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='knowledge', resource_id=knowledge.id, @@ -464,7 +464,7 @@ async def update_knowledge_by_id( detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - form_data.access_grants = filter_allowed_access_grants( + form_data.access_grants = await filter_allowed_access_grants( request.app.state.config.USER_PERMISSIONS, user.id, user.role, @@ -472,7 +472,7 @@ async def update_knowledge_by_id( 'sharing.public_knowledge', ) - knowledge = Knowledges.update_knowledge_by_id(id=id, form_data=form_data) + knowledge = await Knowledges.update_knowledge_by_id(id=id, form_data=form_data) if knowledge: # Re-embed knowledge base for semantic search await embed_knowledge_base_metadata( @@ -483,7 +483,7 @@ async def update_knowledge_by_id( ) return KnowledgeFilesResponse( **knowledge.model_dump(), - files=Knowledges.get_file_metadatas_by_id(knowledge.id), + files=await Knowledges.get_file_metadatas_by_id(knowledge.id), ) else: raise HTTPException( @@ -507,9 +507,9 @@ async def update_knowledge_access_by_id( id: str, form_data: KnowledgeAccessGrantsForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - knowledge = Knowledges.get_knowledge_by_id(id=id, db=db) + knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db) if not knowledge: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -518,7 +518,7 @@ async def update_knowledge_access_by_id( if ( knowledge.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='knowledge', resource_id=knowledge.id, @@ -532,7 +532,7 @@ async def update_knowledge_access_by_id( detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - form_data.access_grants = filter_allowed_access_grants( + form_data.access_grants = await filter_allowed_access_grants( request.app.state.config.USER_PERMISSIONS, user.id, user.role, @@ -540,11 +540,11 @@ async def update_knowledge_access_by_id( 'sharing.public_knowledge', ) - AccessGrants.set_access_grants('knowledge', id, form_data.access_grants, db=db) + await AccessGrants.set_access_grants('knowledge', id, form_data.access_grants, db=db) return KnowledgeFilesResponse( - **Knowledges.get_knowledge_by_id(id=id, db=db).model_dump(), - files=Knowledges.get_file_metadatas_by_id(id, db=db), + **await Knowledges.get_knowledge_by_id(id=id, db=db).model_dump(), + files=await Knowledges.get_file_metadatas_by_id(id, db=db), ) @@ -562,9 +562,9 @@ async def get_knowledge_files_by_id( direction: Optional[str] = None, page: Optional[int] = 1, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - knowledge = Knowledges.get_knowledge_by_id(id=id, db=db) + knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db) if not knowledge: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -574,7 +574,7 @@ async def get_knowledge_files_by_id( if not ( user.role == 'admin' or knowledge.user_id == user.id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='knowledge', resource_id=knowledge.id, @@ -602,7 +602,7 @@ async def get_knowledge_files_by_id( if direction: filter['direction'] = direction - return Knowledges.search_files_by_id(id, user.id, filter=filter, skip=skip, limit=limit, db=db) + return await Knowledges.search_files_by_id(id, user.id, filter=filter, skip=skip, limit=limit, db=db) ############################ @@ -615,14 +615,14 @@ class KnowledgeFileIdForm(BaseModel): @router.post('/{id}/file/add', response_model=Optional[KnowledgeFilesResponse]) -def add_file_to_knowledge_by_id( +async def add_file_to_knowledge_by_id( request: Request, id: str, form_data: KnowledgeFileIdForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - knowledge = Knowledges.get_knowledge_by_id(id=id, db=db) + knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db) if not knowledge: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -631,7 +631,7 @@ def add_file_to_knowledge_by_id( if ( knowledge.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='knowledge', resource_id=knowledge.id, @@ -645,7 +645,7 @@ def add_file_to_knowledge_by_id( detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - file = Files.get_file_by_id(form_data.file_id, db=db) + file = await Files.get_file_by_id(form_data.file_id, db=db) if not file: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -667,7 +667,7 @@ def add_file_to_knowledge_by_id( ) # Add file to knowledge base - Knowledges.add_file_to_knowledge_by_id(knowledge_id=id, file_id=form_data.file_id, user_id=user.id, db=db) + await Knowledges.add_file_to_knowledge_by_id(knowledge_id=id, file_id=form_data.file_id, user_id=user.id, db=db) except Exception as e: log.debug(e) raise HTTPException( @@ -678,7 +678,7 @@ def add_file_to_knowledge_by_id( if knowledge: return KnowledgeFilesResponse( **knowledge.model_dump(), - files=Knowledges.get_file_metadatas_by_id(knowledge.id, db=db), + files=await Knowledges.get_file_metadatas_by_id(knowledge.id, db=db), ) else: raise HTTPException( @@ -688,14 +688,14 @@ def add_file_to_knowledge_by_id( @router.post('/{id}/file/update', response_model=Optional[KnowledgeFilesResponse]) -def update_file_from_knowledge_by_id( +async def update_file_from_knowledge_by_id( request: Request, id: str, form_data: KnowledgeFileIdForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - knowledge = Knowledges.get_knowledge_by_id(id=id, db=db) + knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db) if not knowledge: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -704,7 +704,7 @@ def update_file_from_knowledge_by_id( if ( knowledge.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='knowledge', resource_id=knowledge.id, @@ -718,7 +718,7 @@ def update_file_from_knowledge_by_id( detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - file = Files.get_file_by_id(form_data.file_id, db=db) + file = await Files.get_file_by_id(form_data.file_id, db=db) if not file: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -726,7 +726,7 @@ def update_file_from_knowledge_by_id( ) # Validate the file actually belongs to this knowledge base - if not Knowledges.has_file(knowledge_id=id, file_id=form_data.file_id, db=db): + if not await Knowledges.has_file(knowledge_id=id, file_id=form_data.file_id, db=db): raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.NOT_FOUND, @@ -752,7 +752,7 @@ def update_file_from_knowledge_by_id( if knowledge: return KnowledgeFilesResponse( **knowledge.model_dump(), - files=Knowledges.get_file_metadatas_by_id(knowledge.id, db=db), + files=await Knowledges.get_file_metadatas_by_id(knowledge.id, db=db), ) else: raise HTTPException( @@ -767,14 +767,14 @@ def update_file_from_knowledge_by_id( @router.post('/{id}/file/remove', response_model=Optional[KnowledgeFilesResponse]) -def remove_file_from_knowledge_by_id( +async def remove_file_from_knowledge_by_id( id: str, form_data: KnowledgeFileIdForm, delete_file: bool = Query(True), user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - knowledge = Knowledges.get_knowledge_by_id(id=id, db=db) + knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db) if not knowledge: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -783,7 +783,7 @@ def remove_file_from_knowledge_by_id( if ( knowledge.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='knowledge', resource_id=knowledge.id, @@ -797,7 +797,7 @@ def remove_file_from_knowledge_by_id( detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - file = Files.get_file_by_id(form_data.file_id, db=db) + file = await Files.get_file_by_id(form_data.file_id, db=db) if not file: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -805,13 +805,13 @@ def remove_file_from_knowledge_by_id( ) # Validate the file actually belongs to this knowledge base - if not Knowledges.has_file(knowledge_id=id, file_id=form_data.file_id, db=db): + if not await Knowledges.has_file(knowledge_id=id, file_id=form_data.file_id, db=db): raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.NOT_FOUND, ) - Knowledges.remove_file_from_knowledge_by_id(knowledge_id=id, file_id=form_data.file_id, db=db) + await Knowledges.remove_file_from_knowledge_by_id(knowledge_id=id, file_id=form_data.file_id, db=db) # Remove content from the vector database try: @@ -839,12 +839,12 @@ def remove_file_from_knowledge_by_id( pass # Delete file from database - Files.delete_file_by_id(form_data.file_id, db=db) + await Files.delete_file_by_id(form_data.file_id, db=db) if knowledge: return KnowledgeFilesResponse( **knowledge.model_dump(), - files=Knowledges.get_file_metadatas_by_id(knowledge.id, db=db), + files=await Knowledges.get_file_metadatas_by_id(knowledge.id, db=db), ) else: raise HTTPException( @@ -859,8 +859,8 @@ def remove_file_from_knowledge_by_id( @router.delete('/{id}/delete', response_model=bool) -async def delete_knowledge_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - knowledge = Knowledges.get_knowledge_by_id(id=id, db=db) +async def delete_knowledge_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db) if not knowledge: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -869,7 +869,7 @@ async def delete_knowledge_by_id(id: str, user=Depends(get_verified_user), db: S if ( knowledge.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='knowledge', resource_id=knowledge.id, @@ -886,7 +886,7 @@ async def delete_knowledge_by_id(id: str, user=Depends(get_verified_user), db: S log.info(f'Deleting knowledge base: {id} (name: {knowledge.name})') # Get all models - models = Models.get_all_models(db=db) + models = await Models.get_all_models(db=db) log.info(f'Found {len(models)} models to check for knowledge base {id}') # Update models that reference this knowledge base @@ -910,7 +910,7 @@ async def delete_knowledge_by_id(id: str, user=Depends(get_verified_user), db: S access_grants=model.access_grants, is_active=model.is_active, ) - Models.update_model_by_id(model.id, model_form, db=db) + await Models.update_model_by_id(model.id, model_form, db=db) # Clean up vector DB try: @@ -922,7 +922,7 @@ async def delete_knowledge_by_id(id: str, user=Depends(get_verified_user), db: S # Remove knowledge base embedding remove_knowledge_base_metadata_embedding(id) - result = Knowledges.delete_knowledge_by_id(id=id, db=db) + result = await Knowledges.delete_knowledge_by_id(id=id, db=db) return result @@ -932,8 +932,8 @@ async def delete_knowledge_by_id(id: str, user=Depends(get_verified_user), db: S @router.post('/{id}/reset', response_model=Optional[KnowledgeResponse]) -async def reset_knowledge_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - knowledge = Knowledges.get_knowledge_by_id(id=id, db=db) +async def reset_knowledge_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db) if not knowledge: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -942,7 +942,7 @@ async def reset_knowledge_by_id(id: str, user=Depends(get_verified_user), db: Se if ( knowledge.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='knowledge', resource_id=knowledge.id, @@ -977,12 +977,12 @@ async def add_files_to_knowledge_batch( id: str, form_data: list[KnowledgeFileIdForm], user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """ Add multiple files to a knowledge base """ - knowledge = Knowledges.get_knowledge_by_id(id=id, db=db) + knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db) if not knowledge: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -991,7 +991,7 @@ async def add_files_to_knowledge_batch( if ( knowledge.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='knowledge', resource_id=knowledge.id, @@ -1008,7 +1008,7 @@ async def add_files_to_knowledge_batch( # Batch-fetch all files to avoid N+1 queries log.info(f'files/batch/add - {len(form_data)} files') file_ids = [form.file_id for form in form_data] - files = Files.get_files_by_ids(file_ids, db=db) + files = await Files.get_files_by_ids(file_ids, db=db) # Verify all requested files were found found_ids = {file.id for file in files} @@ -1034,14 +1034,14 @@ async def add_files_to_knowledge_batch( # Only add files that were successfully processed successful_file_ids = [r.file_id for r in result.results if r.status == 'completed'] for file_id in successful_file_ids: - Knowledges.add_file_to_knowledge_by_id(knowledge_id=id, file_id=file_id, user_id=user.id, db=db) + await Knowledges.add_file_to_knowledge_by_id(knowledge_id=id, file_id=file_id, user_id=user.id, db=db) # If there were any errors, include them in the response if result.errors: error_details = [f'{err.file_id}: {err.error}' for err in result.errors] return KnowledgeFilesResponse( **knowledge.model_dump(), - files=Knowledges.get_file_metadatas_by_id(knowledge.id, db=db), + files=await Knowledges.get_file_metadatas_by_id(knowledge.id, db=db), warnings={ 'message': 'Some files failed to process', 'errors': error_details, @@ -1050,7 +1050,7 @@ async def add_files_to_knowledge_batch( return KnowledgeFilesResponse( **knowledge.model_dump(), - files=Knowledges.get_file_metadatas_by_id(knowledge.id, db=db), + files=await Knowledges.get_file_metadatas_by_id(knowledge.id, db=db), ) @@ -1060,20 +1060,20 @@ async def add_files_to_knowledge_batch( @router.get('/{id}/export') -async def export_knowledge_by_id(id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)): +async def export_knowledge_by_id(id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): """ Export a knowledge base as a zip file containing .txt files. Admin only. """ - knowledge = Knowledges.get_knowledge_by_id(id=id, db=db) + knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db) if not knowledge: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND, ) - files = Knowledges.get_files_by_id(id, db=db) + files = await Knowledges.get_files_by_id(id, db=db) # Create zip file in memory zip_buffer = io.BytesIO() diff --git a/backend/open_webui/routers/memories.py b/backend/open_webui/routers/memories.py index 4557f0c44d..3f6c080ad6 100644 --- a/backend/open_webui/routers/memories.py +++ b/backend/open_webui/routers/memories.py @@ -7,8 +7,8 @@ from typing import Optional from open_webui.models.memories import Memories, MemoryModel from open_webui.retrieval.vector.factory import VECTOR_DB_CLIENT from open_webui.utils.auth import get_verified_user -from open_webui.internal.db import get_session -from sqlalchemy.orm import Session +from open_webui.internal.db import get_async_session +from sqlalchemy.ext.asyncio import AsyncSession from open_webui.utils.access_control import has_permission from open_webui.constants import ERROR_MESSAGES @@ -29,7 +29,7 @@ router = APIRouter() async def get_memories( request: Request, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): if not request.app.state.config.ENABLE_MEMORIES: raise HTTPException( @@ -37,13 +37,13 @@ async def get_memories( detail=ERROR_MESSAGES.NOT_FOUND, ) - if not has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS): + if not await has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - return Memories.get_memories_by_user_id(user.id, db=db) + return await Memories.get_memories_by_user_id(user.id, db=db) ############################ @@ -65,7 +65,7 @@ async def add_memory( form_data: AddMemoryForm, user=Depends(get_verified_user), ): - # NOTE: We intentionally do NOT use Depends(get_session) here. + # NOTE: We intentionally do NOT use Depends(get_async_session) here. # Database operations (insert_new_memory) manage their own short-lived sessions. # This prevents holding a connection during EMBEDDING_FUNCTION() # which makes external embedding API calls (1-5+ seconds). @@ -75,13 +75,13 @@ async def add_memory( detail=ERROR_MESSAGES.NOT_FOUND, ) - if not has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS): + if not await has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - memory = Memories.insert_new_memory(user.id, form_data.content) + memory = await Memories.insert_new_memory(user.id, form_data.content) vector = await request.app.state.EMBEDDING_FUNCTION(memory.content, user=user) @@ -116,7 +116,7 @@ async def query_memory( form_data: QueryMemoryForm, user=Depends(get_verified_user), ): - # NOTE: We intentionally do NOT use Depends(get_session) here. + # NOTE: We intentionally do NOT use Depends(get_async_session) here. # Database operations (get_memories_by_user_id) manage their own short-lived sessions. # This prevents holding a connection during EMBEDDING_FUNCTION() # which makes external embedding API calls (1-5+ seconds). @@ -126,13 +126,13 @@ async def query_memory( detail=ERROR_MESSAGES.NOT_FOUND, ) - if not has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS): + if not await has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - memories = Memories.get_memories_by_user_id(user.id) + memories = await Memories.get_memories_by_user_id(user.id) if not memories: raise HTTPException(status_code=404, detail='No memories found for user') @@ -157,7 +157,7 @@ async def reset_memory_from_vector_db( ): """Reset user's memory vector embeddings. - CRITICAL: We intentionally do NOT use Depends(get_session) here. + CRITICAL: We intentionally do NOT use Depends(get_async_session) here. This endpoint generates embeddings for ALL user memories in parallel using asyncio.gather(). A user with 100 memories would trigger 100 embedding API calls simultaneously. With a session held, this could block a connection @@ -169,7 +169,7 @@ async def reset_memory_from_vector_db( detail=ERROR_MESSAGES.NOT_FOUND, ) - if not has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS): + if not await has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.ACCESS_PROHIBITED, @@ -177,7 +177,7 @@ async def reset_memory_from_vector_db( VECTOR_DB_CLIENT.delete_collection(f'user-memory-{user.id}') - memories = Memories.get_memories_by_user_id(user.id) + memories = await Memories.get_memories_by_user_id(user.id) # Generate vectors in parallel vectors = await asyncio.gather( @@ -212,7 +212,7 @@ async def reset_memory_from_vector_db( async def delete_memory_by_user_id( request: Request, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): if not request.app.state.config.ENABLE_MEMORIES: raise HTTPException( @@ -220,13 +220,13 @@ async def delete_memory_by_user_id( detail=ERROR_MESSAGES.NOT_FOUND, ) - if not has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS): + if not await has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - result = Memories.delete_memories_by_user_id(user.id, db=db) + result = await Memories.delete_memories_by_user_id(user.id, db=db) if result: try: @@ -250,7 +250,7 @@ async def update_memory_by_id( form_data: MemoryUpdateModel, user=Depends(get_verified_user), ): - # NOTE: We intentionally do NOT use Depends(get_session) here. + # NOTE: We intentionally do NOT use Depends(get_async_session) here. # Database operations (update_memory_by_id_and_user_id) manage their own # short-lived sessions. This prevents holding a connection during # EMBEDDING_FUNCTION() which makes external API calls (1-5+ seconds). @@ -260,13 +260,13 @@ async def update_memory_by_id( detail=ERROR_MESSAGES.NOT_FOUND, ) - if not has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS): + if not await has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - memory = Memories.update_memory_by_id_and_user_id(memory_id, user.id, form_data.content) + memory = await Memories.update_memory_by_id_and_user_id(memory_id, user.id, form_data.content) if memory is None: raise HTTPException(status_code=404, detail='Memory not found') @@ -301,7 +301,7 @@ async def delete_memory_by_id( memory_id: str, request: Request, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): if not request.app.state.config.ENABLE_MEMORIES: raise HTTPException( @@ -309,13 +309,13 @@ async def delete_memory_by_id( detail=ERROR_MESSAGES.NOT_FOUND, ) - if not has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS): + if not await has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - result = Memories.delete_memory_by_id_and_user_id(memory_id, user.id, db=db) + result = await Memories.delete_memory_by_id_and_user_id(memory_id, user.id, db=db) if result: VECTOR_DB_CLIENT.delete(collection_name=f'user-memory-{user.id}', ids=[memory_id]) diff --git a/backend/open_webui/routers/models.py b/backend/open_webui/routers/models.py index 6f7b3d48df..b7d31321ff 100644 --- a/backend/open_webui/routers/models.py +++ b/backend/open_webui/routers/models.py @@ -35,8 +35,8 @@ from fastapi.responses import FileResponse, StreamingResponse from open_webui.utils.auth import get_admin_user, get_verified_user from open_webui.utils.access_control import has_permission, filter_allowed_access_grants from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL, STATIC_DIR -from open_webui.internal.db import get_session -from sqlalchemy.orm import Session +from open_webui.internal.db import get_async_session +from sqlalchemy.ext.asyncio import AsyncSession log = logging.getLogger(__name__) @@ -66,7 +66,7 @@ async def get_models( direction: Optional[str] = None, page: Optional[int] = 1, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): limit = PAGE_ITEM_COUNT @@ -86,7 +86,7 @@ async def get_models( filter['direction'] = direction # Pre-fetch user group IDs once - used for both filter and write_access check - groups = Groups.get_groups_by_member_id(user.id, db=db) + groups = await Groups.get_groups_by_member_id(user.id, db=db) user_group_ids = {group.id for group in groups} if not user.role == 'admin' or not BYPASS_ADMIN_ACCESS_CONTROL: @@ -95,11 +95,11 @@ async def get_models( filter['user_id'] = user.id - result = Models.search_models(user.id, filter=filter, skip=skip, limit=limit, db=db) + result = await Models.search_models(user.id, filter=filter, skip=skip, limit=limit, db=db) # Batch-fetch writable model IDs in a single query instead of N has_access calls model_ids = [model.id for model in result.items] - writable_model_ids = AccessGrants.get_accessible_resource_ids( + writable_model_ids = await AccessGrants.get_accessible_resource_ids( user_id=user.id, resource_type='model', resource_ids=model_ids, @@ -130,8 +130,8 @@ async def get_models( @router.get('/base', response_model=list[ModelResponse]) -async def get_base_models(user=Depends(get_admin_user), db: Session = Depends(get_session)): - return Models.get_base_models(db=db) +async def get_base_models(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): + return await Models.get_base_models(db=db) ########################### @@ -140,11 +140,11 @@ async def get_base_models(user=Depends(get_admin_user), db: Session = Depends(ge @router.get('/tags', response_model=list[str]) -async def get_model_tags(user=Depends(get_verified_user), db: Session = Depends(get_session)): +async def get_model_tags(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): if user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL: - models = Models.get_models(db=db) + models = await Models.get_models(db=db) else: - models = Models.get_models_by_user_id(user.id, db=db) + models = await Models.get_models_by_user_id(user.id, db=db) tags_set = set() for model in models: @@ -172,9 +172,9 @@ async def create_new_model( request: Request, form_data: ModelForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - if user.role != 'admin' and not has_permission( + if user.role != 'admin' and not await has_permission( user.id, 'workspace.models', request.app.state.config.USER_PERMISSIONS, db=db ): raise HTTPException( @@ -182,7 +182,7 @@ async def create_new_model( detail=ERROR_MESSAGES.UNAUTHORIZED, ) - model = Models.get_model_by_id(form_data.id, db=db) + model = await Models.get_model_by_id(form_data.id, db=db) if model: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, @@ -196,7 +196,7 @@ async def create_new_model( ) else: - form_data.access_grants = filter_allowed_access_grants( + form_data.access_grants = await filter_allowed_access_grants( request.app.state.config.USER_PERMISSIONS, user.id, user.role, @@ -204,7 +204,7 @@ async def create_new_model( 'sharing.public_models', ) - model = Models.insert_new_model(form_data, user.id, db=db) + model = await Models.insert_new_model(form_data, user.id, db=db) if model: return model else: @@ -223,9 +223,9 @@ async def create_new_model( async def export_models( request: Request, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - if user.role != 'admin' and not has_permission( + if user.role != 'admin' and not await has_permission( user.id, 'workspace.models_export', request.app.state.config.USER_PERMISSIONS, @@ -237,9 +237,9 @@ async def export_models( ) if user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL: - return Models.get_models(db=db) + return await Models.get_models(db=db) else: - return Models.get_models_by_user_id(user.id, db=db) + return await Models.get_models_by_user_id(user.id, db=db) ############################ @@ -256,9 +256,9 @@ async def import_models( request: Request, user=Depends(get_verified_user), form_data: ModelsImportForm = (...), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - if user.role != 'admin' and not has_permission( + if user.role != 'admin' and not await has_permission( user.id, 'workspace.models_import', request.app.state.config.USER_PERMISSIONS, @@ -278,7 +278,7 @@ async def import_models( if model_data.get('id') and is_valid_model_id(model_data.get('id')) ] existing_models = { - model.id: model for model in (Models.get_models_by_ids(model_ids, db=db) if model_ids else []) + model.id: model for model in (await Models.get_models_by_ids(model_ids, db=db) if model_ids else []) } for model_data in data: @@ -293,13 +293,13 @@ async def import_models( model_data['params'] = model_data.get('params', {}) updated_model = ModelForm(**{**existing_model.model_dump(), **model_data}) - Models.update_model_by_id(model_id, updated_model, db=db) + await Models.update_model_by_id(model_id, updated_model, db=db) else: # Insert new model model_data['meta'] = model_data.get('meta', {}) model_data['params'] = model_data.get('params', {}) new_model = ModelForm(**model_data) - Models.insert_new_model(user_id=user.id, form_data=new_model, db=db) + await Models.insert_new_model(user_id=user.id, form_data=new_model, db=db) return True else: raise HTTPException(status_code=400, detail='Invalid JSON format') @@ -322,9 +322,9 @@ async def sync_models( request: Request, form_data: SyncModelsForm, user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - return Models.sync_models(user.id, form_data.models, db=db) + return await Models.sync_models(user.id, form_data.models, db=db) ########################### @@ -338,13 +338,13 @@ class ModelIdForm(BaseModel): # Note: We're not using the typical url path param here, but instead using a query parameter to allow '/' in the id @router.get('/model', response_model=Optional[ModelAccessResponse]) -async def get_model_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - model = Models.get_model_by_id(id, db=db) +async def get_model_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + model = await Models.get_model_by_id(id, db=db) if model: if ( (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL) or model.user_id == user.id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='model', resource_id=model.id, @@ -357,7 +357,7 @@ async def get_model_by_id(id: str, user=Depends(get_verified_user), db: Session write_access=( (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL) or user.id == model.user_id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='model', resource_id=model.id, @@ -384,8 +384,8 @@ async def get_model_by_id(id: str, user=Depends(get_verified_user), db: Session @router.get('/model/profile/image') -def get_model_profile_image(id: str, user=Depends(get_verified_user)): - model = Models.get_model_by_id(id) +async def get_model_profile_image(id: str, user=Depends(get_verified_user)): + model = await Models.get_model_by_id(id) if model: etag = f'"{model.updated_at}"' if model.updated_at else None @@ -426,13 +426,13 @@ def get_model_profile_image(id: str, user=Depends(get_verified_user)): @router.post('/model/toggle', response_model=Optional[ModelResponse]) -async def toggle_model_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - model = Models.get_model_by_id(id, db=db) +async def toggle_model_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + model = await Models.get_model_by_id(id, db=db) if model: if ( user.role == 'admin' or model.user_id == user.id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='model', resource_id=model.id, @@ -440,7 +440,7 @@ async def toggle_model_by_id(id: str, user=Depends(get_verified_user), db: Sessi db=db, ) ): - model = Models.toggle_model_by_id(id, db=db) + model = await Models.toggle_model_by_id(id, db=db) if model: return model @@ -471,9 +471,9 @@ async def update_model_by_id( request: Request, form_data: ModelForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - model = Models.get_model_by_id(form_data.id, db=db) + model = await Models.get_model_by_id(form_data.id, db=db) if not model: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, @@ -482,7 +482,7 @@ async def update_model_by_id( if ( model.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='model', resource_id=model.id, @@ -496,7 +496,7 @@ async def update_model_by_id( detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - form_data.access_grants = filter_allowed_access_grants( + form_data.access_grants = await filter_allowed_access_grants( request.app.state.config.USER_PERMISSIONS, user.id, user.role, @@ -504,7 +504,7 @@ async def update_model_by_id( 'sharing.public_models', ) - model = Models.update_model_by_id(form_data.id, ModelForm(**form_data.model_dump()), db=db) + model = await Models.update_model_by_id(form_data.id, ModelForm(**form_data.model_dump()), db=db) return model @@ -524,9 +524,9 @@ async def update_model_access_by_id( request: Request, form_data: ModelAccessGrantsForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - model = Models.get_model_by_id(form_data.id, db=db) + model = await Models.get_model_by_id(form_data.id, db=db) # Non-preset models (e.g. direct Ollama/OpenAI models) may not have a DB # entry yet. Create a minimal one so access grants can be stored. @@ -536,7 +536,7 @@ async def update_model_access_by_id( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - model = Models.insert_new_model( + model = await Models.insert_new_model( ModelForm( id=form_data.id, name=form_data.name or form_data.id, @@ -554,7 +554,7 @@ async def update_model_access_by_id( if ( model.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='model', resource_id=model.id, @@ -568,7 +568,7 @@ async def update_model_access_by_id( detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - form_data.access_grants = filter_allowed_access_grants( + form_data.access_grants = await filter_allowed_access_grants( request.app.state.config.USER_PERMISSIONS, user.id, user.role, @@ -576,11 +576,11 @@ async def update_model_access_by_id( 'sharing.public_models', ) - AccessGrants.set_access_grants('model', form_data.id, form_data.access_grants, db=db) + await AccessGrants.set_access_grants('model', form_data.id, form_data.access_grants, db=db) - Models.update_model_updated_at_by_id(form_data.id, db=db) + await Models.update_model_updated_at_by_id(form_data.id, db=db) - return Models.get_model_by_id(form_data.id, db=db) + return await Models.get_model_by_id(form_data.id, db=db) ############################ @@ -592,9 +592,9 @@ async def update_model_access_by_id( async def delete_model_by_id( form_data: ModelIdForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - model = Models.get_model_by_id(form_data.id, db=db) + model = await Models.get_model_by_id(form_data.id, db=db) if not model: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, @@ -604,7 +604,7 @@ async def delete_model_by_id( if ( user.role != 'admin' and model.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='model', resource_id=model.id, @@ -617,11 +617,11 @@ async def delete_model_by_id( detail=ERROR_MESSAGES.UNAUTHORIZED, ) - result = Models.delete_model_by_id(form_data.id, db=db) + result = await Models.delete_model_by_id(form_data.id, db=db) return result @router.delete('/delete/all', response_model=bool) -async def delete_all_models(user=Depends(get_admin_user), db: Session = Depends(get_session)): - result = Models.delete_all_models(db=db) +async def delete_all_models(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): + result = await Models.delete_all_models(db=db) return result diff --git a/backend/open_webui/routers/notes.py b/backend/open_webui/routers/notes.py index 0eec88a251..61c9fb7d95 100644 --- a/backend/open_webui/routers/notes.py +++ b/backend/open_webui/routers/notes.py @@ -34,8 +34,8 @@ from open_webui.utils.access_control import ( filter_allowed_access_grants, ) from open_webui.models.access_grants import AccessGrants -from open_webui.internal.db import get_session -from sqlalchemy.orm import Session +from open_webui.internal.db import get_async_session +from sqlalchemy.ext.asyncio import AsyncSession log = logging.getLogger(__name__) @@ -68,9 +68,9 @@ async def get_notes( request: Request, page: Optional[int] = None, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - if user.role != 'admin' and not has_permission( + if user.role != 'admin' and not await has_permission( user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db ): raise HTTPException( @@ -84,12 +84,12 @@ async def get_notes( limit = 60 skip = (page - 1) * limit - notes = Notes.get_notes_by_user_id(user.id, 'read', skip=skip, limit=limit, db=db) + notes = await Notes.get_notes_by_user_id(user.id, 'read', skip=skip, limit=limit, db=db) if not notes: return [] user_ids = list(set(note.user_id for note in notes)) - users = {user.id: user for user in Users.get_users_by_user_ids(user_ids, db=db)} + users = {user.id: user for user in await Users.get_users_by_user_ids(user_ids, db=db)} return [ NoteUserResponse( @@ -114,9 +114,9 @@ async def search_notes( direction: Optional[str] = None, page: Optional[int] = 1, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - if user.role != 'admin' and not has_permission( + if user.role != 'admin' and not await has_permission( user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db ): raise HTTPException( @@ -143,13 +143,13 @@ async def search_notes( filter['direction'] = direction if not user.role == 'admin' or not BYPASS_ADMIN_ACCESS_CONTROL: - groups = Groups.get_groups_by_member_id(user.id, db=db) + groups = await Groups.get_groups_by_member_id(user.id, db=db) if groups: filter['group_ids'] = [group.id for group in groups] filter['user_id'] = user.id - result = Notes.search_notes(user.id, filter, skip=skip, limit=limit, db=db) + result = await Notes.search_notes(user.id, filter, skip=skip, limit=limit, db=db) for note in result.items: note.data = _truncate_note_data(note.data) return result @@ -165,9 +165,9 @@ async def create_new_note( request: Request, form_data: NoteForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - if user.role != 'admin' and not has_permission( + if user.role != 'admin' and not await has_permission( user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db ): raise HTTPException( @@ -175,7 +175,7 @@ async def create_new_note( detail=ERROR_MESSAGES.UNAUTHORIZED, ) - form_data.access_grants = filter_allowed_access_grants( + form_data.access_grants = await filter_allowed_access_grants( request.app.state.config.USER_PERMISSIONS, user.id, user.role, @@ -185,7 +185,7 @@ async def create_new_note( ) try: - note = Notes.insert_new_note(user.id, form_data, db=db) + note = await Notes.insert_new_note(user.id, form_data, db=db) return note except Exception as e: log.exception(e) @@ -206,9 +206,9 @@ async def get_note_by_id( request: Request, id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - if user.role != 'admin' and not has_permission( + if user.role != 'admin' and not await has_permission( user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db ): raise HTTPException( @@ -216,14 +216,14 @@ async def get_note_by_id( detail=ERROR_MESSAGES.UNAUTHORIZED, ) - note = Notes.get_note_by_id(id, db=db) + note = await Notes.get_note_by_id(id, db=db) if not note: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) if user.role != 'admin' and ( user.id != note.user_id and ( - not AccessGrants.has_access( + not await AccessGrants.has_access( user_id=user.id, resource_type='note', resource_id=note.id, @@ -237,7 +237,7 @@ async def get_note_by_id( write_access = ( user.role == 'admin' or (user.id == note.user_id) - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='note', resource_id=note.id, @@ -261,9 +261,9 @@ async def update_note_by_id( id: str, form_data: NoteForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - if user.role != 'admin' and not has_permission( + if user.role != 'admin' and not await has_permission( user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db ): raise HTTPException( @@ -271,13 +271,13 @@ async def update_note_by_id( detail=ERROR_MESSAGES.UNAUTHORIZED, ) - note = Notes.get_note_by_id(id, db=db) + note = await Notes.get_note_by_id(id, db=db) if not note: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) if user.role != 'admin' and ( user.id != note.user_id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='note', resource_id=note.id, @@ -287,7 +287,7 @@ async def update_note_by_id( ): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) - form_data.access_grants = filter_allowed_access_grants( + form_data.access_grants = await filter_allowed_access_grants( request.app.state.config.USER_PERMISSIONS, user.id, user.role, @@ -297,7 +297,7 @@ async def update_note_by_id( ) try: - note = Notes.update_note_by_id(id, form_data, db=db) + note = await Notes.update_note_by_id(id, form_data, db=db) await sio.emit( 'note-events', note.model_dump(), @@ -325,9 +325,9 @@ async def update_note_access_by_id( id: str, form_data: NoteAccessGrantsForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - if user.role != 'admin' and not has_permission( + if user.role != 'admin' and not await has_permission( user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db ): raise HTTPException( @@ -335,13 +335,13 @@ async def update_note_access_by_id( detail=ERROR_MESSAGES.UNAUTHORIZED, ) - note = Notes.get_note_by_id(id, db=db) + note = await Notes.get_note_by_id(id, db=db) if not note: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) if user.role != 'admin' and ( user.id != note.user_id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='note', resource_id=note.id, @@ -351,7 +351,7 @@ async def update_note_access_by_id( ): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) - form_data.access_grants = filter_allowed_access_grants( + form_data.access_grants = await filter_allowed_access_grants( request.app.state.config.USER_PERMISSIONS, user.id, user.role, @@ -359,9 +359,9 @@ async def update_note_access_by_id( 'sharing.public_notes', ) - AccessGrants.set_access_grants('note', id, form_data.access_grants, db=db) + await AccessGrants.set_access_grants('note', id, form_data.access_grants, db=db) - return Notes.get_note_by_id(id, db=db) + return await Notes.get_note_by_id(id, db=db) ############################ @@ -374,9 +374,9 @@ async def delete_note_by_id( request: Request, id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - if user.role != 'admin' and not has_permission( + if user.role != 'admin' and not await has_permission( user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db ): raise HTTPException( @@ -384,13 +384,13 @@ async def delete_note_by_id( detail=ERROR_MESSAGES.UNAUTHORIZED, ) - note = Notes.get_note_by_id(id, db=db) + note = await Notes.get_note_by_id(id, db=db) if not note: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) if user.role != 'admin' and ( user.id != note.user_id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='note', resource_id=note.id, @@ -401,7 +401,7 @@ async def delete_note_by_id( raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) try: - note = Notes.delete_note_by_id(id, db=db) + note = await Notes.delete_note_by_id(id, db=db) return True except Exception as e: log.exception(e) diff --git a/backend/open_webui/routers/ollama.py b/backend/open_webui/routers/ollama.py index 5cf87854ac..11c916846a 100644 --- a/backend/open_webui/routers/ollama.py +++ b/backend/open_webui/routers/ollama.py @@ -39,9 +39,9 @@ from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import StreamingResponse from pydantic import BaseModel, ConfigDict, validator -from sqlalchemy.orm import Session +from sqlalchemy.ext.asyncio import AsyncSession -from open_webui.internal.db import get_session +from open_webui.internal.db import get_async_session from open_webui.models.models import Models @@ -398,11 +398,11 @@ async def get_all_models(request: Request, user: UserModel = None): async def get_filtered_models(models, user, db=None): # Filter models based on user access control model_ids = [model['model'] for model in models.get('models', [])] - model_infos = {model_info.id: model_info for model_info in Models.get_models_by_ids(model_ids, db=db)} - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id, db=db)} + model_infos = {model_info.id: model_info for model_info in await Models.get_models_by_ids(model_ids, db=db)} + user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)} # Batch-fetch accessible resource IDs in a single query instead of N has_access calls - accessible_model_ids = AccessGrants.get_accessible_resource_ids( + accessible_model_ids = await AccessGrants.get_accessible_resource_ids( user_id=user.id, resource_type='model', resource_ids=list(model_infos.keys()), @@ -800,7 +800,7 @@ async def show_model_info(request: Request, form_data: ModelNameForm, user=Depen model = form_data.get('model') # Enforce per-model access control - check_model_access(user, Models.get_model_by_id(model), BYPASS_MODEL_ACCESS_CONTROL) + await check_model_access(user, await Models.get_model_by_id(model), BYPASS_MODEL_ACCESS_CONTROL) await get_all_models(request, user=user) models = request.app.state.OLLAMA_MODELS @@ -850,7 +850,7 @@ async def embed( log.info(f'generate_ollama_batch_embeddings {form_data}') # Enforce per-model access control - check_model_access(user, Models.get_model_by_id(form_data.model), BYPASS_MODEL_ACCESS_CONTROL) + await check_model_access(user, await Models.get_model_by_id(form_data.model), BYPASS_MODEL_ACCESS_CONTROL) if url_idx is None: model = form_data.model @@ -909,7 +909,7 @@ async def embeddings( log.info(f'generate_ollama_embeddings {form_data}') # Enforce per-model access control - check_model_access(user, Models.get_model_by_id(form_data.model), BYPASS_MODEL_ACCESS_CONTROL) + await check_model_access(user, await Models.get_model_by_id(form_data.model), BYPASS_MODEL_ACCESS_CONTROL) if url_idx is None: model = form_data.model @@ -974,7 +974,7 @@ async def generate_completion( raise HTTPException(status_code=503, detail='Ollama API is disabled') # Enforce per-model access control - check_model_access(user, Models.get_model_by_id(form_data.model), BYPASS_MODEL_ACCESS_CONTROL) + await check_model_access(user, await Models.get_model_by_id(form_data.model), BYPASS_MODEL_ACCESS_CONTROL) if url_idx is None: await get_all_models(request, user=user) @@ -1064,7 +1064,7 @@ async def generate_chat_completion( if not request.app.state.config.ENABLE_OLLAMA_API: raise HTTPException(status_code=503, detail='Ollama API is disabled') - # NOTE: We intentionally do NOT use Depends(get_session) here. + # NOTE: We intentionally do NOT use Depends(get_async_session) here. # Database operations (get_model_by_id, AccessGrants.has_access) manage their own short-lived sessions. # This prevents holding a connection during the entire LLM call (30-60+ seconds), # which would exhaust the connection pool under concurrent load. @@ -1093,7 +1093,7 @@ async def generate_chat_completion( del payload['metadata'] model_id = payload['model'] - model_info = Models.get_model_by_id(model_id) + model_info = await Models.get_model_by_id(model_id) if model_info: if model_info.base_model_id: @@ -1111,9 +1111,9 @@ async def generate_chat_completion( if not bypass_system_prompt: payload = apply_system_prompt_to_body(system, payload, metadata, user) - check_model_access(user, model_info, bypass_filter) + await check_model_access(user, model_info, bypass_filter) else: - check_model_access(user, None, bypass_filter) + await check_model_access(user, None, bypass_filter) url, url_idx = await get_ollama_url(request, payload['model'], url_idx) api_config = request.app.state.config.OLLAMA_API_CONFIGS.get( @@ -1171,7 +1171,7 @@ async def generate_openai_completion( url_idx: Optional[int] = None, user=Depends(get_verified_user), ): - # NOTE: We intentionally do NOT use Depends(get_session) here. + # NOTE: We intentionally do NOT use Depends(get_async_session) here. # Database operations (get_model_by_id, AccessGrants.has_access) manage their own short-lived sessions. # This prevents holding a connection during the entire LLM call (30-60+ seconds), # which would exhaust the connection pool under concurrent load. @@ -1191,7 +1191,7 @@ async def generate_openai_completion( del payload['metadata'] model_id = form_data.model - model_info = Models.get_model_by_id(model_id) + model_info = await Models.get_model_by_id(model_id) if model_info: if model_info.base_model_id: payload['model'] = model_info.base_model_id @@ -1200,9 +1200,9 @@ async def generate_openai_completion( if params: payload = apply_model_params_to_body_openai(params, payload) - check_model_access(user, model_info) + await check_model_access(user, model_info) else: - check_model_access(user, None) + await check_model_access(user, None) url, url_idx = await get_ollama_url(request, payload['model'], url_idx) api_config = request.app.state.config.OLLAMA_API_CONFIGS.get( @@ -1233,7 +1233,7 @@ async def generate_openai_chat_completion( url_idx: Optional[int] = None, user=Depends(get_verified_user), ): - # NOTE: We intentionally do NOT use Depends(get_session) here. + # NOTE: We intentionally do NOT use Depends(get_async_session) here. # Database operations (get_model_by_id, AccessGrants.has_access) manage their own short-lived sessions. # This prevents holding a connection during the entire LLM call (30-60+ seconds), # which would exhaust the connection pool under concurrent load. @@ -1253,7 +1253,7 @@ async def generate_openai_chat_completion( del payload['metadata'] model_id = completion_form.model - model_info = Models.get_model_by_id(model_id) + model_info = await Models.get_model_by_id(model_id) if model_info: if model_info.base_model_id: payload['model'] = model_info.base_model_id @@ -1266,9 +1266,9 @@ async def generate_openai_chat_completion( payload = apply_model_params_to_body_openai(params, payload) payload = apply_system_prompt_to_body(system, payload, metadata, user) - check_model_access(user, model_info) + await check_model_access(user, model_info) else: - check_model_access(user, None) + await check_model_access(user, None) url, url_idx = await get_ollama_url(request, payload['model'], url_idx) api_config = request.app.state.config.OLLAMA_API_CONFIGS.get( @@ -1313,14 +1313,14 @@ async def generate_anthropic_messages( payload = {**form_data} model_id = payload.get('model', '') - model_info = Models.get_model_by_id(model_id) + model_info = await Models.get_model_by_id(model_id) if model_info: if model_info.base_model_id: payload['model'] = model_info.base_model_id - check_model_access(user, model_info) + await check_model_access(user, model_info) else: - check_model_access(user, None) + await check_model_access(user, None) url, url_idx = await get_ollama_url(request, payload['model'], url_idx) api_config = request.app.state.config.OLLAMA_API_CONFIGS.get( @@ -1371,17 +1371,17 @@ async def generate_responses( payload = form_data.model_dump() model_id = form_data.model - model_info = Models.get_model_by_id(model_id) + model_info = await Models.get_model_by_id(model_id) if model_info: if model_info.base_model_id: payload['model'] = model_info.base_model_id # Check if user has access to the model if user.role == 'user': - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)} + user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)} if not ( user.id == model_info.user_id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='model', resource_id=model_info.id, @@ -1426,7 +1426,7 @@ async def get_openai_models( request: Request, url_idx: Optional[int] = None, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): models = [] if url_idx is None: @@ -1458,11 +1458,11 @@ async def get_openai_models( if user.role == 'user' and not BYPASS_MODEL_ACCESS_CONTROL: # Filter models based on user access control model_ids = [model['id'] for model in models] - model_infos = {model_info.id: model_info for model_info in Models.get_models_by_ids(model_ids, db=db)} - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id, db=db)} + model_infos = {model_info.id: model_info for model_info in await Models.get_models_by_ids(model_ids, db=db)} + user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)} # Batch-fetch accessible resource IDs in a single query instead of N has_access calls - accessible_model_ids = AccessGrants.get_accessible_resource_ids( + accessible_model_ids = await AccessGrants.get_accessible_resource_ids( user_id=user.id, resource_type='model', resource_ids=list(model_infos.keys()), diff --git a/backend/open_webui/routers/openai.py b/backend/open_webui/routers/openai.py index 14225a243f..047e56fcf6 100644 --- a/backend/open_webui/routers/openai.py +++ b/backend/open_webui/routers/openai.py @@ -21,9 +21,9 @@ from fastapi.responses import ( ) from pydantic import BaseModel, ConfigDict -from sqlalchemy.orm import Session +from sqlalchemy.ext.asyncio import AsyncSession -from open_webui.internal.db import get_session +from open_webui.internal.db import get_async_session from open_webui.models.models import Models from open_webui.models.access_grants import AccessGrants @@ -450,11 +450,11 @@ async def get_all_models_responses(request: Request, user: UserModel) -> list: async def get_filtered_models(models, user, db=None): # Filter models based on user access control model_ids = [model['id'] for model in models.get('data', [])] - model_infos = {model_info.id: model_info for model_info in Models.get_models_by_ids(model_ids, db=db)} - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id, db=db)} + model_infos = {model_info.id: model_info for model_info in await Models.get_models_by_ids(model_ids, db=db)} + user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)} # Batch-fetch accessible resource IDs in a single query instead of N has_access calls - accessible_model_ids = AccessGrants.get_accessible_resource_ids( + accessible_model_ids = await AccessGrants.get_accessible_resource_ids( user_id=user.id, resource_type='model', resource_ids=list(model_infos.keys()), @@ -1025,7 +1025,7 @@ async def generate_chat_completion( user=Depends(get_verified_user), bypass_system_prompt: bool = False, ): - # NOTE: We intentionally do NOT use Depends(get_session) here. + # NOTE: We intentionally do NOT use Depends(get_async_session) here. # Database operations (get_model_by_id, AccessGrants.has_access) manage their own short-lived sessions. # This prevents holding a connection during the entire LLM call (30-60+ seconds), # which would exhaust the connection pool under concurrent load. @@ -1043,7 +1043,7 @@ async def generate_chat_completion( metadata = payload.pop('metadata', None) model_id = form_data.get('model') - model_info = Models.get_model_by_id(model_id) + model_info = await Models.get_model_by_id(model_id) # Check model info and override the payload if model_info: @@ -1063,9 +1063,9 @@ async def generate_chat_completion( if not bypass_system_prompt: payload = apply_system_prompt_to_body(system, payload, metadata, user) - check_model_access(user, model_info, bypass_filter) + await check_model_access(user, model_info, bypass_filter) else: - check_model_access(user, None, bypass_filter) + await check_model_access(user, None, bypass_filter) # Check if model is already in app state cache to avoid expensive get_all_models() call models = request.app.state.OPENAI_MODELS @@ -1344,7 +1344,7 @@ async def responses( model_id = form_data.model # Enforce per-model access control - check_model_access(user, Models.get_model_by_id(model_id), BYPASS_MODEL_ACCESS_CONTROL) + await check_model_access(user, await Models.get_model_by_id(model_id), BYPASS_MODEL_ACCESS_CONTROL) body = json.dumps(payload) diff --git a/backend/open_webui/routers/prompts.py b/backend/open_webui/routers/prompts.py index 3b579c2892..ed9b69af06 100644 --- a/backend/open_webui/routers/prompts.py +++ b/backend/open_webui/routers/prompts.py @@ -20,8 +20,8 @@ from open_webui.constants import ERROR_MESSAGES from open_webui.utils.auth import get_admin_user, get_verified_user from open_webui.utils.access_control import has_permission, filter_allowed_access_grants from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL -from open_webui.internal.db import get_session -from sqlalchemy.orm import Session +from open_webui.internal.db import get_async_session +from sqlalchemy.ext.asyncio import AsyncSession from pydantic import BaseModel @@ -48,21 +48,21 @@ PAGE_ITEM_COUNT = 30 @router.get('/', response_model=list[PromptModel]) -async def get_prompts(user=Depends(get_verified_user), db: Session = Depends(get_session)): +async def get_prompts(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): if user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL: - prompts = Prompts.get_prompts(db=db) + prompts = await Prompts.get_prompts(db=db) else: - prompts = Prompts.get_prompts_by_user_id(user.id, 'read', db=db) + prompts = await Prompts.get_prompts_by_user_id(user.id, 'read', db=db) return prompts @router.get('/tags', response_model=list[str]) -async def get_prompt_tags(user=Depends(get_verified_user), db: Session = Depends(get_session)): +async def get_prompt_tags(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): if user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL: - return Prompts.get_tags(db=db) + return await Prompts.get_tags(db=db) else: - prompts = Prompts.get_prompts_by_user_id(user.id, 'read', db=db) + prompts = await Prompts.get_prompts_by_user_id(user.id, 'read', db=db) tags = set() for prompt in prompts: if prompt.tags: @@ -79,7 +79,7 @@ async def get_prompt_list( direction: Optional[str] = None, page: Optional[int] = 1, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): limit = PAGE_ITEM_COUNT @@ -99,7 +99,7 @@ async def get_prompt_list( filter['direction'] = direction # Pre-fetch user group IDs once - used for both filter and write_access check - groups = Groups.get_groups_by_member_id(user.id, db=db) + groups = await Groups.get_groups_by_member_id(user.id, db=db) user_group_ids = {group.id for group in groups} if not (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL): @@ -108,11 +108,11 @@ async def get_prompt_list( filter['user_id'] = user.id - result = Prompts.search_prompts(user.id, filter=filter, skip=skip, limit=limit, db=db) + result = await Prompts.search_prompts(user.id, filter=filter, skip=skip, limit=limit, db=db) # Batch-fetch writable prompt IDs in a single query instead of N has_access calls prompt_ids = [prompt.id for prompt in result.items] - writable_prompt_ids = AccessGrants.get_accessible_resource_ids( + writable_prompt_ids = await AccessGrants.get_accessible_resource_ids( user_id=user.id, resource_type='prompt', resource_ids=prompt_ids, @@ -147,16 +147,16 @@ async def create_new_prompt( request: Request, form_data: PromptForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): if user.role != 'admin' and not ( - has_permission( + await has_permission( user.id, 'workspace.prompts', request.app.state.config.USER_PERMISSIONS, db=db, ) - or has_permission( + or await has_permission( user.id, 'workspace.prompts_import', request.app.state.config.USER_PERMISSIONS, @@ -168,7 +168,7 @@ async def create_new_prompt( detail=ERROR_MESSAGES.UNAUTHORIZED, ) - form_data.access_grants = filter_allowed_access_grants( + form_data.access_grants = await filter_allowed_access_grants( request.app.state.config.USER_PERMISSIONS, user.id, user.role, @@ -176,9 +176,9 @@ async def create_new_prompt( 'sharing.public_prompts', ) - prompt = Prompts.get_prompt_by_command(form_data.command, db=db) + prompt = await Prompts.get_prompt_by_command(form_data.command, db=db) if prompt is None: - prompt = Prompts.insert_new_prompt(user.id, form_data, db=db) + prompt = await Prompts.insert_new_prompt(user.id, form_data, db=db) if prompt: return prompt @@ -198,14 +198,14 @@ async def create_new_prompt( @router.get('/command/{command}', response_model=Optional[PromptAccessResponse]) -async def get_prompt_by_command(command: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - prompt = Prompts.get_prompt_by_command(command, db=db) +async def get_prompt_by_command(command: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + prompt = await Prompts.get_prompt_by_command(command, db=db) if prompt: if ( user.role == 'admin' or prompt.user_id == user.id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='prompt', resource_id=prompt.id, @@ -218,7 +218,7 @@ async def get_prompt_by_command(command: str, user=Depends(get_verified_user), d write_access=( (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL) or user.id == prompt.user_id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='prompt', resource_id=prompt.id, @@ -240,14 +240,14 @@ async def get_prompt_by_command(command: str, user=Depends(get_verified_user), d @router.get('/id/{prompt_id}', response_model=Optional[PromptAccessResponse]) -async def get_prompt_by_id(prompt_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - prompt = Prompts.get_prompt_by_id(prompt_id, db=db) +async def get_prompt_by_id(prompt_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + prompt = await Prompts.get_prompt_by_id(prompt_id, db=db) if prompt: if ( user.role == 'admin' or prompt.user_id == user.id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='prompt', resource_id=prompt.id, @@ -260,7 +260,7 @@ async def get_prompt_by_id(prompt_id: str, user=Depends(get_verified_user), db: write_access=( (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL) or user.id == prompt.user_id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='prompt', resource_id=prompt.id, @@ -287,9 +287,9 @@ async def update_prompt_by_id( prompt_id: str, form_data: PromptForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - prompt = Prompts.get_prompt_by_id(prompt_id, db=db) + prompt = await Prompts.get_prompt_by_id(prompt_id, db=db) if not prompt: raise HTTPException( @@ -300,7 +300,7 @@ async def update_prompt_by_id( # Is the user the original creator, in a group with write access, or an admin if ( prompt.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='prompt', resource_id=prompt.id, @@ -316,14 +316,14 @@ async def update_prompt_by_id( # Check for command collision if command is being changed if form_data.command != prompt.command: - existing_prompt = Prompts.get_prompt_by_command(form_data.command, db=db) + existing_prompt = await Prompts.get_prompt_by_command(form_data.command, db=db) if existing_prompt and existing_prompt.id != prompt.id: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, detail=f"Command '/{form_data.command}' is already in use by another prompt", ) - form_data.access_grants = filter_allowed_access_grants( + form_data.access_grants = await filter_allowed_access_grants( request.app.state.config.USER_PERMISSIONS, user.id, user.role, @@ -332,7 +332,7 @@ async def update_prompt_by_id( ) # Use the ID from the found prompt - updated_prompt = Prompts.update_prompt_by_id(prompt.id, form_data, user.id, db=db) + updated_prompt = await Prompts.update_prompt_by_id(prompt.id, form_data, user.id, db=db) if updated_prompt: return updated_prompt else: @@ -352,10 +352,10 @@ async def update_prompt_metadata( prompt_id: str, form_data: PromptMetadataForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Update prompt name and command only (no history created).""" - prompt = Prompts.get_prompt_by_id(prompt_id, db=db) + prompt = await Prompts.get_prompt_by_id(prompt_id, db=db) if not prompt: raise HTTPException( @@ -365,7 +365,7 @@ async def update_prompt_metadata( if ( prompt.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='prompt', resource_id=prompt.id, @@ -381,14 +381,14 @@ async def update_prompt_metadata( # Check for command collision if command is being changed if form_data.command != prompt.command: - existing_prompt = Prompts.get_prompt_by_command(form_data.command, db=db) + existing_prompt = await Prompts.get_prompt_by_command(form_data.command, db=db) if existing_prompt and existing_prompt.id != prompt.id: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, detail=f"Command '/{form_data.command}' is already in use", ) - updated_prompt = Prompts.update_prompt_metadata(prompt.id, form_data.name, form_data.command, form_data.tags, db=db) + updated_prompt = await Prompts.update_prompt_metadata(prompt.id, form_data.name, form_data.command, form_data.tags, db=db) if updated_prompt: return updated_prompt else: @@ -403,9 +403,9 @@ async def set_prompt_version( prompt_id: str, form_data: PromptVersionUpdateForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - prompt = Prompts.get_prompt_by_id(prompt_id, db=db) + prompt = await Prompts.get_prompt_by_id(prompt_id, db=db) if not prompt: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -414,7 +414,7 @@ async def set_prompt_version( if ( prompt.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='prompt', resource_id=prompt.id, @@ -428,7 +428,7 @@ async def set_prompt_version( detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - updated_prompt = Prompts.update_prompt_version(prompt.id, form_data.version_id, db=db) + updated_prompt = await Prompts.update_prompt_version(prompt.id, form_data.version_id, db=db) if updated_prompt: return updated_prompt else: @@ -453,9 +453,9 @@ async def update_prompt_access_by_id( prompt_id: str, form_data: PromptAccessGrantsForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - prompt = Prompts.get_prompt_by_id(prompt_id, db=db) + prompt = await Prompts.get_prompt_by_id(prompt_id, db=db) if not prompt: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -464,7 +464,7 @@ async def update_prompt_access_by_id( if ( prompt.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='prompt', resource_id=prompt.id, @@ -478,7 +478,7 @@ async def update_prompt_access_by_id( detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - form_data.access_grants = filter_allowed_access_grants( + form_data.access_grants = await filter_allowed_access_grants( request.app.state.config.USER_PERMISSIONS, user.id, user.role, @@ -486,9 +486,9 @@ async def update_prompt_access_by_id( 'sharing.public_prompts', ) - AccessGrants.set_access_grants('prompt', prompt_id, form_data.access_grants, db=db) + await AccessGrants.set_access_grants('prompt', prompt_id, form_data.access_grants, db=db) - return Prompts.get_prompt_by_id(prompt_id, db=db) + return await Prompts.get_prompt_by_id(prompt_id, db=db) ############################ @@ -497,8 +497,8 @@ async def update_prompt_access_by_id( @router.post('/id/{prompt_id}/toggle', response_model=Optional[PromptModel]) -async def toggle_prompt_active(prompt_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - prompt = Prompts.get_prompt_by_id(prompt_id, db=db) +async def toggle_prompt_active(prompt_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + prompt = await Prompts.get_prompt_by_id(prompt_id, db=db) if not prompt: raise HTTPException( @@ -508,7 +508,7 @@ async def toggle_prompt_active(prompt_id: str, user=Depends(get_verified_user), if ( prompt.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='prompt', resource_id=prompt.id, @@ -522,7 +522,7 @@ async def toggle_prompt_active(prompt_id: str, user=Depends(get_verified_user), detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - result = Prompts.toggle_prompt_active(prompt.id, db=db) + result = await Prompts.toggle_prompt_active(prompt.id, db=db) if result: return result raise HTTPException( @@ -537,8 +537,8 @@ async def toggle_prompt_active(prompt_id: str, user=Depends(get_verified_user), @router.delete('/id/{prompt_id}/delete', response_model=bool) -async def delete_prompt_by_id(prompt_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - prompt = Prompts.get_prompt_by_id(prompt_id, db=db) +async def delete_prompt_by_id(prompt_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + prompt = await Prompts.get_prompt_by_id(prompt_id, db=db) if not prompt: raise HTTPException( @@ -548,7 +548,7 @@ async def delete_prompt_by_id(prompt_id: str, user=Depends(get_verified_user), d if ( prompt.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='prompt', resource_id=prompt.id, @@ -562,7 +562,7 @@ async def delete_prompt_by_id(prompt_id: str, user=Depends(get_verified_user), d detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - result = Prompts.delete_prompt_by_id(prompt.id, db=db) + result = await Prompts.delete_prompt_by_id(prompt.id, db=db) return result @@ -576,12 +576,12 @@ async def get_prompt_history( prompt_id: str, page: int = 0, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Get version history for a prompt.""" PAGE_SIZE = 20 - prompt = Prompts.get_prompt_by_id(prompt_id, db=db) + prompt = await Prompts.get_prompt_by_id(prompt_id, db=db) if not prompt: raise HTTPException( @@ -593,7 +593,7 @@ async def get_prompt_history( if not ( user.role == 'admin' or prompt.user_id == user.id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='prompt', resource_id=prompt.id, @@ -606,7 +606,7 @@ async def get_prompt_history( detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - history = PromptHistories.get_history_by_prompt_id(prompt.id, limit=PAGE_SIZE, offset=page * PAGE_SIZE, db=db) + history = await PromptHistories.get_history_by_prompt_id(prompt.id, limit=PAGE_SIZE, offset=page * PAGE_SIZE, db=db) return history @@ -615,10 +615,10 @@ async def get_prompt_history_entry( prompt_id: str, history_id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Get a specific version from history.""" - prompt = Prompts.get_prompt_by_id(prompt_id, db=db) + prompt = await Prompts.get_prompt_by_id(prompt_id, db=db) if not prompt: raise HTTPException( @@ -630,7 +630,7 @@ async def get_prompt_history_entry( if not ( user.role == 'admin' or prompt.user_id == user.id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='prompt', resource_id=prompt.id, @@ -643,7 +643,7 @@ async def get_prompt_history_entry( detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - history_entry = PromptHistories.get_history_entry_by_id(history_id, db=db) + history_entry = await PromptHistories.get_history_entry_by_id(history_id, db=db) if not history_entry or history_entry.prompt_id != prompt.id: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -658,10 +658,10 @@ async def delete_prompt_history_entry( prompt_id: str, history_id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Delete a history entry. Cannot delete the active production version.""" - prompt = Prompts.get_prompt_by_id(prompt_id, db=db) + prompt = await Prompts.get_prompt_by_id(prompt_id, db=db) if not prompt: raise HTTPException( @@ -673,7 +673,7 @@ async def delete_prompt_history_entry( if not ( user.role == 'admin' or prompt.user_id == user.id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='prompt', resource_id=prompt.id, @@ -693,7 +693,7 @@ async def delete_prompt_history_entry( detail='Cannot delete the active production version', ) - success = PromptHistories.delete_history_entry(history_id, db=db) + success = await PromptHistories.delete_history_entry(history_id, db=db) if not success: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -709,10 +709,10 @@ async def get_prompt_diff( from_id: str, to_id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Get diff between two versions.""" - prompt = Prompts.get_prompt_by_id(prompt_id, db=db) + prompt = await Prompts.get_prompt_by_id(prompt_id, db=db) if not prompt: raise HTTPException( @@ -724,7 +724,7 @@ async def get_prompt_diff( if not ( user.role == 'admin' or prompt.user_id == user.id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='prompt', resource_id=prompt.id, @@ -737,7 +737,7 @@ async def get_prompt_diff( detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - diff = PromptHistories.compute_diff(from_id, to_id, db=db) + diff = await PromptHistories.compute_diff(from_id, to_id, db=db) if not diff: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, diff --git a/backend/open_webui/routers/retrieval.py b/backend/open_webui/routers/retrieval.py index 0fa971684f..77261a6f2b 100644 --- a/backend/open_webui/routers/retrieval.py +++ b/backend/open_webui/routers/retrieval.py @@ -40,8 +40,8 @@ from open_webui.models.files import FileModel, FileUpdateForm, Files from open_webui.utils.access_control.files import has_access_to_file from open_webui.models.knowledge import Knowledges from open_webui.storage.provider import Storage -from open_webui.internal.db import get_session, get_db -from sqlalchemy.orm import Session +from open_webui.internal.db import get_async_session, get_db +from sqlalchemy.ext.asyncio import AsyncSession from open_webui.retrieval.vector.factory import VECTOR_DB_CLIENT @@ -1526,11 +1526,11 @@ class ProcessFileForm(BaseModel): @router.post('/process/file') -def process_file( +async def process_file( request: Request, form_data: ProcessFileForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """ Process a file and save its content to the vector database. @@ -1539,9 +1539,9 @@ def process_file( The session is committed before external API calls, and updates use a fresh session. """ if user.role == 'admin': - file = Files.get_file_by_id(form_data.file_id, db=db) + file = await Files.get_file_by_id(form_data.file_id, db=db) else: - file = Files.get_file_by_id_and_user_id(form_data.file_id, user.id, db=db) + file = await Files.get_file_by_id_and_user_id(form_data.file_id, user.id, db=db) if file: try: @@ -1674,7 +1674,7 @@ def process_file( text_content = ' '.join([doc.page_content for doc in docs]) log.debug(f'text_content: {text_content}') - Files.update_file_data_by_id( + await Files.update_file_data_by_id( file.id, {'content': text_content}, db=db, @@ -1682,8 +1682,8 @@ def process_file( hash = calculate_sha256_string(text_content) if request.app.state.config.BYPASS_EMBEDDING_AND_RETRIEVAL: - Files.update_file_data_by_id(file.id, {'status': 'completed'}, db=db) - Files.update_file_hash_by_id(file.id, hash, db=db) + await Files.update_file_data_by_id(file.id, {'status': 'completed'}, db=db) + await Files.update_file_hash_by_id(file.id, hash, db=db) return { 'status': True, 'collection_name': None, @@ -1715,7 +1715,7 @@ def process_file( if result: # Fresh session for the final update. with get_db() as session: - Files.update_file_metadata_by_id( + await Files.update_file_metadata_by_id( file.id, { 'collection_name': collection_name, @@ -1723,12 +1723,12 @@ def process_file( db=session, ) - Files.update_file_data_by_id( + await Files.update_file_data_by_id( file.id, {'status': 'completed'}, db=session, ) - Files.update_file_hash_by_id(file.id, hash, db=session) + await Files.update_file_hash_by_id(file.id, hash, db=session) return { 'status': True, @@ -1745,13 +1745,13 @@ def process_file( log.exception(e) # Fresh session for error status update. with get_db() as session: - Files.update_file_data_by_id( + await Files.update_file_data_by_id( file.id, {'status': 'failed'}, db=session, ) # Clear the hash so the file can be re-uploaded after fixing the issue - Files.update_file_hash_by_id(file.id, None, db=session) + await Files.update_file_hash_by_id(file.id, None, db=session) if 'No pandoc was found' in str(e): raise HTTPException( @@ -2175,7 +2175,7 @@ async def process_web_search(request: Request, form_data: SearchForm, user=Depen detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - if user.role != 'admin' and not has_permission( + if user.role != 'admin' and not await has_permission( user.id, 'features.web_search', request.app.state.config.USER_PERMISSIONS ): raise HTTPException( @@ -2327,7 +2327,7 @@ async def process_web_search(request: Request, form_data: SearchForm, user=Depen ) -def _validate_collection_access(collection_names: list[str], user) -> None: +async def _validate_collection_access(collection_names: list[str], user) -> None: """ Prevent users from querying collections they don't own. Enforces ownership on user-memory-* and file-* collections. @@ -2344,7 +2344,7 @@ def _validate_collection_access(collection_names: list[str], user) -> None: ) elif name.startswith('file-'): file_id = name[len('file-') :] - if not has_access_to_file( + if not await has_access_to_file( file_id=file_id, access_type='read', user=user, @@ -2370,7 +2370,7 @@ async def query_doc_handler( form_data: QueryDocForm, user=Depends(get_verified_user), ): - _validate_collection_access([form_data.collection_name], user) + await _validate_collection_access([form_data.collection_name], user) try: if request.app.state.config.ENABLE_RAG_HYBRID_SEARCH and (form_data.hybrid is None or form_data.hybrid): @@ -2435,7 +2435,7 @@ async def query_collection_handler( form_data: QueryCollectionsForm, user=Depends(get_verified_user), ): - _validate_collection_access(form_data.collection_names, user) + await _validate_collection_access(form_data.collection_names, user) try: if request.app.state.config.ENABLE_RAG_HYBRID_SEARCH and (form_data.hybrid is None or form_data.hybrid): @@ -2496,14 +2496,14 @@ class DeleteForm(BaseModel): @router.post('/delete') -def delete_entries_from_collection( +async def delete_entries_from_collection( form_data: DeleteForm, user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): try: if VECTOR_DB_CLIENT.has_collection(collection_name=form_data.collection_name): - file = Files.get_file_by_id(form_data.file_id, db=db) + file = await Files.get_file_by_id(form_data.file_id, db=db) if not file: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -2524,13 +2524,13 @@ def delete_entries_from_collection( @router.post('/reset/db') -def reset_vector_db(user=Depends(get_admin_user), db: Session = Depends(get_session)): +async def reset_vector_db(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): VECTOR_DB_CLIENT.reset() - Knowledges.delete_all_knowledge(db=db) + await Knowledges.delete_all_knowledge(db=db) @router.post('/reset/uploads') -def reset_upload_dir(user=Depends(get_admin_user)) -> bool: +async def reset_upload_dir(user=Depends(get_admin_user)) -> bool: folder = f'{UPLOAD_DIR}' try: # Check if the directory exists @@ -2585,7 +2585,7 @@ async def process_files_batch( """ Process a batch of files and save them to the vector database. - NOTE: We intentionally do NOT use Depends(get_session) here. + NOTE: We intentionally do NOT use Depends(get_async_session) here. The save_docs_to_vector_db() call makes external embedding API calls which can take 5-60+ seconds for batch operations. Database operations after embedding (Files.update_file_by_id) manage their own short-lived sessions. @@ -2603,7 +2603,7 @@ async def process_files_batch( for file in form_data.files: try: # Ownership check: verify the requesting user owns the file or is an admin - db_file = Files.get_file_by_id(file.id, db=db) + db_file = await Files.get_file_by_id(file.id, db=db) if not db_file: file_errors.append( BatchProcessFilesResult( @@ -2665,7 +2665,7 @@ async def process_files_batch( # Update all files with collection name for file_update, file_result in zip(file_updates, file_results): - Files.update_file_by_id(id=file_result.file_id, form_data=file_update, db=db) + await Files.update_file_by_id(id=file_result.file_id, form_data=file_update, db=db) file_result.status = 'completed' except Exception as e: diff --git a/backend/open_webui/routers/scim.py b/backend/open_webui/routers/scim.py index 7bc0157b19..75f45bcaf9 100644 --- a/backend/open_webui/routers/scim.py +++ b/backend/open_webui/routers/scim.py @@ -30,8 +30,8 @@ from open_webui.config import OAUTH_PROVIDERS from open_webui.env import SCIM_AUTH_PROVIDER -from sqlalchemy.orm import Session -from open_webui.internal.db import get_session +from sqlalchemy.ext.asyncio import AsyncSession +from open_webui.internal.db import get_async_session log = logging.getLogger(__name__) @@ -326,18 +326,18 @@ def get_scim_provider() -> str: return SCIM_AUTH_PROVIDER -def find_user_by_external_id(external_id: str, db=None) -> Optional[UserModel]: +async def find_user_by_external_id(external_id: str, db=None) -> Optional[UserModel]: """Find a user by SCIM externalId, falling back to OAuth sub match.""" provider = get_scim_provider() - user = Users.get_user_by_scim_external_id(provider, external_id, db=db) + user = await Users.get_user_by_scim_external_id(provider, external_id, db=db) if user: return user # Fallback: check if externalId matches an existing OAuth sub (account linking) - return Users.get_user_by_oauth_sub(provider, external_id, db=db) + return await Users.get_user_by_oauth_sub(provider, external_id, db=db) -def user_to_scim(user: UserModel, request: Request, db=None) -> SCIMUser: +async def user_to_scim(user: UserModel, request: Request, db=None) -> SCIMUser: """Convert internal User model to SCIM User""" # Parse display name into name components name_parts = user.name.split(' ', 1) if user.name else ['', ''] @@ -345,7 +345,7 @@ def user_to_scim(user: UserModel, request: Request, db=None) -> SCIMUser: family_name = name_parts[1] if len(name_parts) > 1 else '' # Get user's groups - user_groups = Groups.get_groups_by_member_id(user.id, db=db) + user_groups = await Groups.get_groups_by_member_id(user.id, db=db) groups = [ { 'value': group.id, @@ -379,12 +379,12 @@ def user_to_scim(user: UserModel, request: Request, db=None) -> SCIMUser: ) -def group_to_scim(group: GroupModel, request: Request, db=None) -> SCIMGroup: +async def group_to_scim(group: GroupModel, request: Request, db=None) -> SCIMGroup: """Convert internal Group model to SCIM Group""" - member_ids = Groups.get_group_user_ids_by_id(group.id, db) or [] + member_ids = await Groups.get_group_user_ids_by_id(group.id, db) or [] # Batch-fetch all users to avoid N+1 queries - users = Users.get_users_by_user_ids(member_ids, db=db) if member_ids else [] + users = await Users.get_users_by_user_ids(member_ids, db=db) if member_ids else [] members = [ SCIMGroupMember( value=user.id, @@ -512,7 +512,7 @@ async def get_users( count: int = Query(20), filter: Optional[str] = None, _: bool = Depends(get_scim_auth), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """List SCIM Users""" # Clamp per SCIM 2.0 spec (RFC 7644 §3.4.2.4): @@ -527,25 +527,25 @@ async def get_users( # Simple filter parsing - supports userName eq, externalId eq if 'userName eq' in filter: email = filter.split('"')[1] - user = Users.get_user_by_email(email, db=db) + user = await Users.get_user_by_email(email, db=db) users_list = [user] if user else [] total = 1 if user else 0 elif 'externalId eq' in filter: external_id = filter.split('"')[1] - user = find_user_by_external_id(external_id, db=db) + user = await find_user_by_external_id(external_id, db=db) users_list = [user] if user else [] total = 1 if user else 0 else: - response = Users.get_users(skip=skip, limit=limit, db=db) + response = await Users.get_users(skip=skip, limit=limit, db=db) users_list = response['users'] total = response['total'] else: - response = Users.get_users(skip=skip, limit=limit, db=db) + response = await Users.get_users(skip=skip, limit=limit, db=db) users_list = response['users'] total = response['total'] # Convert to SCIM format - scim_users = [user_to_scim(user, request, db=db) for user in users_list] + scim_users = [await user_to_scim(user, request, db=db) for user in users_list] return SCIMListResponse( totalResults=total, @@ -560,14 +560,14 @@ async def get_user( user_id: str, request: Request, _: bool = Depends(get_scim_auth), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Get SCIM User by ID""" - user = Users.get_user_by_id(user_id, db=db) + user = await Users.get_user_by_id(user_id, db=db) if not user: return scim_error(status_code=status.HTTP_404_NOT_FOUND, detail=f'User {user_id} not found') - return user_to_scim(user, request, db=db) + return await user_to_scim(user, request, db=db) @router.post('/Users', response_model=SCIMUser, status_code=status.HTTP_201_CREATED) @@ -575,12 +575,12 @@ async def create_user( request: Request, user_data: SCIMUserCreateRequest, _: bool = Depends(get_scim_auth), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Create SCIM User""" # Check for duplicate by externalId if user_data.externalId: - existing_user = find_user_by_external_id(user_data.externalId, db=db) + existing_user = await find_user_by_external_id(user_data.externalId, db=db) if existing_user: raise HTTPException( status_code=status.HTTP_409_CONFLICT, @@ -596,7 +596,7 @@ async def create_user( email = email.lower() # Check for duplicate by email - existing_user = Users.get_user_by_email(email, db=db) + existing_user = await Users.get_user_by_email(email, db=db) if existing_user: raise HTTPException( status_code=status.HTTP_409_CONFLICT, @@ -619,7 +619,7 @@ async def create_user( if user_data.photos and len(user_data.photos) > 0: profile_image = user_data.photos[0].value - new_user = Users.insert_new_user( + new_user = await Users.insert_new_user( id=user_id, name=name, email=email, @@ -637,10 +637,10 @@ async def create_user( # Store externalId in the scim field if user_data.externalId: provider = get_scim_provider() - Users.update_user_scim_by_id(user_id, provider, user_data.externalId, db=db) - new_user = Users.get_user_by_id(user_id, db=db) + await Users.update_user_scim_by_id(user_id, provider, user_data.externalId, db=db) + new_user = await Users.get_user_by_id(user_id, db=db) - return user_to_scim(new_user, request, db=db) + return await user_to_scim(new_user, request, db=db) @router.put('/Users/{user_id}', response_model=SCIMUser) @@ -649,10 +649,10 @@ async def update_user( request: Request, user_data: SCIMUserUpdateRequest, _: bool = Depends(get_scim_auth), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Update SCIM User (full update)""" - user = Users.get_user_by_id(user_id, db=db) + user = await Users.get_user_by_id(user_id, db=db) if not user: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -682,7 +682,7 @@ async def update_user( if user_data.photos and len(user_data.photos) > 0: update_data['profile_image_url'] = user_data.photos[0].value - updated_user = Users.update_user_by_id(user_id, update_data, db=db) + updated_user = await Users.update_user_by_id(user_id, update_data, db=db) if not updated_user: raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, @@ -692,10 +692,10 @@ async def update_user( # Update externalId in the scim field if user_data.externalId: provider = get_scim_provider() - Users.update_user_scim_by_id(user_id, provider, user_data.externalId, db=db) - updated_user = Users.get_user_by_id(user_id, db=db) + await Users.update_user_scim_by_id(user_id, provider, user_data.externalId, db=db) + updated_user = await Users.get_user_by_id(user_id, db=db) - return user_to_scim(updated_user, request, db=db) + return await user_to_scim(updated_user, request, db=db) @router.patch('/Users/{user_id}', response_model=SCIMUser) @@ -704,10 +704,10 @@ async def patch_user( request: Request, patch_data: SCIMPatchRequest, _: bool = Depends(get_scim_auth), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Update SCIM User (partial update)""" - user = Users.get_user_by_id(user_id, db=db) + user = await Users.get_user_by_id(user_id, db=db) if not user: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -734,11 +734,11 @@ async def patch_user( update_data['name'] = value elif path == 'externalId': provider = get_scim_provider() - Users.update_user_scim_by_id(user_id, provider, value, db=db) + await Users.update_user_scim_by_id(user_id, provider, value, db=db) # Update user if update_data: - updated_user = Users.update_user_by_id(user_id, update_data, db=db) + updated_user = await Users.update_user_by_id(user_id, update_data, db=db) if not updated_user: raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, @@ -747,7 +747,7 @@ async def patch_user( else: updated_user = user - return user_to_scim(updated_user, request, db=db) + return await user_to_scim(updated_user, request, db=db) @router.delete('/Users/{user_id}', status_code=status.HTTP_204_NO_CONTENT) @@ -755,17 +755,17 @@ async def delete_user( user_id: str, request: Request, _: bool = Depends(get_scim_auth), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Delete SCIM User""" - user = Users.get_user_by_id(user_id, db=db) + user = await Users.get_user_by_id(user_id, db=db) if not user: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, detail=f'User {user_id} not found', ) - success = Users.delete_user_by_id(user_id, db=db) + success = await Users.delete_user_by_id(user_id, db=db) if not success: raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, @@ -783,7 +783,7 @@ async def get_groups( count: int = Query(20), filter: Optional[str] = None, _: bool = Depends(get_scim_auth), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """List SCIM Groups""" # Clamp per SCIM 2.0 spec (RFC 7644 §3.4.2.4): @@ -795,13 +795,13 @@ async def get_groups( if filter: if 'displayName eq' in filter: display_name = filter.split('"')[1] - group = Groups.get_group_by_name(display_name, db=db) + group = await Groups.get_group_by_name(display_name, db=db) groups_list = [group] if group else [] else: # Unrecognized filter — fall back to all groups - groups_list = Groups.get_all_groups(db=db) + groups_list = await Groups.get_all_groups(db=db) else: - groups_list = Groups.get_all_groups(db=db) + groups_list = await Groups.get_all_groups(db=db) # Apply pagination total = len(groups_list) @@ -810,7 +810,7 @@ async def get_groups( paginated_groups = groups_list[start:end] # Convert to SCIM format - scim_groups = [group_to_scim(group, request, db=db) for group in paginated_groups] + scim_groups = [await group_to_scim(group, request, db=db) for group in paginated_groups] return SCIMListResponse( totalResults=total, @@ -825,17 +825,17 @@ async def get_group( group_id: str, request: Request, _: bool = Depends(get_scim_auth), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Get SCIM Group by ID""" - group = Groups.get_group_by_id(group_id, db=db) + group = await Groups.get_group_by_id(group_id, db=db) if not group: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, detail=f'Group {group_id} not found', ) - return group_to_scim(group, request, db=db) + return await group_to_scim(group, request, db=db) @router.post('/Groups', response_model=SCIMGroup, status_code=status.HTTP_201_CREATED) @@ -843,7 +843,7 @@ async def create_group( request: Request, group_data: SCIMGroupCreateRequest, _: bool = Depends(get_scim_auth), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Create SCIM Group""" # Extract member IDs @@ -861,14 +861,14 @@ async def create_group( ) # Need to get the creating user's ID - we'll use the first admin - admin_user = Users.get_super_admin_user(db=db) + admin_user = await Users.get_super_admin_user(db=db) if not admin_user: raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail='No admin user found', ) - new_group = Groups.insert_new_group(admin_user.id, form, db=db) + new_group = await Groups.insert_new_group(admin_user.id, form, db=db) if not new_group: raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, @@ -884,12 +884,12 @@ async def create_group( description=new_group.description, ) - Groups.update_group_by_id(new_group.id, update_form, db=db) - Groups.set_group_user_ids_by_id(new_group.id, member_ids, db=db) + await Groups.update_group_by_id(new_group.id, update_form, db=db) + await Groups.set_group_user_ids_by_id(new_group.id, member_ids, db=db) - new_group = Groups.get_group_by_id(new_group.id, db=db) + new_group = await Groups.get_group_by_id(new_group.id, db=db) - return group_to_scim(new_group, request, db=db) + return await group_to_scim(new_group, request, db=db) @router.put('/Groups/{group_id}', response_model=SCIMGroup) @@ -898,10 +898,10 @@ async def update_group( request: Request, group_data: SCIMGroupUpdateRequest, _: bool = Depends(get_scim_auth), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Update SCIM Group (full update)""" - group = Groups.get_group_by_id(group_id, db=db) + group = await Groups.get_group_by_id(group_id, db=db) if not group: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -919,17 +919,17 @@ async def update_group( # Handle members if provided if group_data.members is not None: member_ids = [member.value for member in group_data.members] - Groups.set_group_user_ids_by_id(group_id, member_ids, db=db) + await Groups.set_group_user_ids_by_id(group_id, member_ids, db=db) # Update group - updated_group = Groups.update_group_by_id(group_id, update_form, db=db) + updated_group = await Groups.update_group_by_id(group_id, update_form, db=db) if not updated_group: raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail='Failed to update group', ) - return group_to_scim(updated_group, request, db=db) + return await group_to_scim(updated_group, request, db=db) @router.patch('/Groups/{group_id}', response_model=SCIMGroup) @@ -938,10 +938,10 @@ async def patch_group( request: Request, patch_data: SCIMPatchRequest, _: bool = Depends(get_scim_auth), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Update SCIM Group (partial update)""" - group = Groups.get_group_by_id(group_id, db=db) + group = await Groups.get_group_by_id(group_id, db=db) if not group: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -965,7 +965,7 @@ async def patch_group( update_form.name = value elif path == 'members': # Replace all members - Groups.set_group_user_ids_by_id(group_id, [member['value'] for member in value], db=db) + await Groups.set_group_user_ids_by_id(group_id, [member['value'] for member in value], db=db) elif op == 'add': if path == 'members': @@ -973,22 +973,22 @@ async def patch_group( if isinstance(value, list): for member in value: if isinstance(member, dict) and 'value' in member: - Groups.add_users_to_group(group_id, [member['value']], db=db) + await Groups.add_users_to_group(group_id, [member['value']], db=db) elif op == 'remove': if path and path.startswith('members[value eq'): # Remove specific member member_id = path.split('"')[1] - Groups.remove_users_from_group(group_id, [member_id], db=db) + await Groups.remove_users_from_group(group_id, [member_id], db=db) # Update group - updated_group = Groups.update_group_by_id(group_id, update_form, db=db) + updated_group = await Groups.update_group_by_id(group_id, update_form, db=db) if not updated_group: raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail='Failed to update group', ) - return group_to_scim(updated_group, request, db=db) + return await group_to_scim(updated_group, request, db=db) @router.delete('/Groups/{group_id}', status_code=status.HTTP_204_NO_CONTENT) @@ -996,17 +996,17 @@ async def delete_group( group_id: str, request: Request, _: bool = Depends(get_scim_auth), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): """Delete SCIM Group""" - group = Groups.get_group_by_id(group_id, db=db) + group = await Groups.get_group_by_id(group_id, db=db) if not group: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, detail=f'Group {group_id} not found', ) - success = Groups.delete_group_by_id(group_id, db=db) + success = await Groups.delete_group_by_id(group_id, db=db) if not success: raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, diff --git a/backend/open_webui/routers/skills.py b/backend/open_webui/routers/skills.py index 1838914e4a..490d1706d5 100644 --- a/backend/open_webui/routers/skills.py +++ b/backend/open_webui/routers/skills.py @@ -5,9 +5,9 @@ from open_webui.models.groups import Groups from pydantic import BaseModel from fastapi import APIRouter, Depends, HTTPException, Request, status -from sqlalchemy.orm import Session +from sqlalchemy.ext.asyncio import AsyncSession -from open_webui.internal.db import get_session +from open_webui.internal.db import get_async_session from open_webui.models.skills import ( SkillForm, SkillModel, @@ -40,18 +40,18 @@ router = APIRouter() async def get_skills( request: Request, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): if user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL: - skills = Skills.get_skills(db=db) + skills = await Skills.get_skills(db=db) else: - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id, db=db)} - all_skills = Skills.get_skills(db=db) + user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)} + all_skills = await Skills.get_skills(db=db) skills = [ skill for skill in all_skills if skill.user_id == user.id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='skill', resource_id=skill.id, @@ -75,7 +75,7 @@ async def get_skill_list( view_option: Optional[str] = None, page: Optional[int] = 1, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): limit = PAGE_ITEM_COUNT @@ -89,13 +89,13 @@ async def get_skill_list( filter['view_option'] = view_option if not (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL): - groups = Groups.get_groups_by_member_id(user.id, db=db) + groups = await Groups.get_groups_by_member_id(user.id, db=db) if groups: filter['group_ids'] = [group.id for group in groups] filter['user_id'] = user.id - result = Skills.search_skills(user.id, filter=filter, skip=skip, limit=limit, db=db) + result = await Skills.search_skills(user.id, filter=filter, skip=skip, limit=limit, db=db) return SkillAccessListResponse( items=[ @@ -104,7 +104,7 @@ async def get_skill_list( write_access=( (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL) or user.id == skill.user_id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='skill', resource_id=skill.id, @@ -128,9 +128,9 @@ async def get_skill_list( async def export_skills( request: Request, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - if user.role != 'admin' and not has_permission( + if user.role != 'admin' and not await has_permission( user.id, 'workspace.skills', request.app.state.config.USER_PERMISSIONS, @@ -142,9 +142,9 @@ async def export_skills( ) if user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL: - return Skills.get_skills(db=db) + return await Skills.get_skills(db=db) else: - return Skills.get_skills_by_user_id(user.id, 'read', db=db) + return await Skills.get_skills_by_user_id(user.id, 'read', db=db) ############################ @@ -157,9 +157,9 @@ async def create_new_skill( request: Request, form_data: SkillForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - if user.role != 'admin' and not has_permission( + if user.role != 'admin' and not await has_permission( user.id, 'workspace.skills', request.app.state.config.USER_PERMISSIONS, db=db ): raise HTTPException( @@ -169,7 +169,7 @@ async def create_new_skill( form_data.id = form_data.id.lower().replace(' ', '-') - existing = Skills.get_skill_by_id(form_data.id, db=db) + existing = await Skills.get_skill_by_id(form_data.id, db=db) if existing is not None: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -177,7 +177,7 @@ async def create_new_skill( ) try: - skill = Skills.insert_new_skill(user.id, form_data, db=db) + skill = await Skills.insert_new_skill(user.id, form_data, db=db) if skill: return skill else: @@ -199,14 +199,14 @@ async def create_new_skill( @router.get('/id/{id}', response_model=Optional[SkillAccessResponse]) -async def get_skill_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - skill = Skills.get_skill_by_id(id, db=db) +async def get_skill_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + skill = await Skills.get_skill_by_id(id, db=db) if skill: if ( user.role == 'admin' or skill.user_id == user.id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='skill', resource_id=skill.id, @@ -219,7 +219,7 @@ async def get_skill_by_id(id: str, user=Depends(get_verified_user), db: Session write_access=( (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL) or user.id == skill.user_id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='skill', resource_id=skill.id, @@ -251,9 +251,9 @@ async def update_skill_by_id( id: str, form_data: SkillForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - skill = Skills.get_skill_by_id(id, db=db) + skill = await Skills.get_skill_by_id(id, db=db) if not skill: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -262,7 +262,7 @@ async def update_skill_by_id( if ( skill.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='skill', resource_id=skill.id, @@ -281,7 +281,7 @@ async def update_skill_by_id( **form_data.model_dump(exclude={'id'}), } - skill = Skills.update_skill_by_id(id, updated, db=db) + skill = await Skills.update_skill_by_id(id, updated, db=db) if skill: return skill @@ -312,9 +312,9 @@ async def update_skill_access_by_id( id: str, form_data: SkillAccessGrantsForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - skill = Skills.get_skill_by_id(id, db=db) + skill = await Skills.get_skill_by_id(id, db=db) if not skill: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -323,7 +323,7 @@ async def update_skill_access_by_id( if ( skill.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='skill', resource_id=skill.id, @@ -337,7 +337,7 @@ async def update_skill_access_by_id( detail=ERROR_MESSAGES.UNAUTHORIZED, ) - form_data.access_grants = filter_allowed_access_grants( + form_data.access_grants = await filter_allowed_access_grants( request.app.state.config.USER_PERMISSIONS, user.id, user.role, @@ -345,9 +345,9 @@ async def update_skill_access_by_id( 'sharing.public_skills', ) - AccessGrants.set_access_grants('skill', id, form_data.access_grants, db=db) + await AccessGrants.set_access_grants('skill', id, form_data.access_grants, db=db) - return Skills.get_skill_by_id(id, db=db) + return await Skills.get_skill_by_id(id, db=db) ############################ @@ -356,13 +356,13 @@ async def update_skill_access_by_id( @router.post('/id/{id}/toggle', response_model=Optional[SkillModel]) -async def toggle_skill_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - skill = Skills.get_skill_by_id(id, db=db) +async def toggle_skill_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + skill = await Skills.get_skill_by_id(id, db=db) if skill: if ( user.role == 'admin' or skill.user_id == user.id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='skill', resource_id=skill.id, @@ -370,7 +370,7 @@ async def toggle_skill_by_id(id: str, user=Depends(get_verified_user), db: Sessi db=db, ) ): - skill = Skills.toggle_skill_by_id(id, db=db) + skill = await Skills.toggle_skill_by_id(id, db=db) if skill: return skill @@ -401,9 +401,9 @@ async def delete_skill_by_id( request: Request, id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - skill = Skills.get_skill_by_id(id, db=db) + skill = await Skills.get_skill_by_id(id, db=db) if not skill: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -412,7 +412,7 @@ async def delete_skill_by_id( if ( skill.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='skill', resource_id=skill.id, @@ -426,5 +426,5 @@ async def delete_skill_by_id( detail=ERROR_MESSAGES.UNAUTHORIZED, ) - result = Skills.delete_skill_by_id(id, db=db) + result = await Skills.delete_skill_by_id(id, db=db) return result diff --git a/backend/open_webui/routers/terminals.py b/backend/open_webui/routers/terminals.py index 34d5eb96d6..0d607d1f78 100644 --- a/backend/open_webui/routers/terminals.py +++ b/backend/open_webui/routers/terminals.py @@ -52,7 +52,7 @@ def _sanitize_proxy_path(path: str) -> str | None: async def list_terminal_servers(request: Request, user=Depends(get_verified_user)): """Return terminal servers the authenticated user has access to.""" connections = request.app.state.config.TERMINAL_SERVER_CONNECTIONS or [] - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)} + user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)} return [ { @@ -61,7 +61,7 @@ async def list_terminal_servers(request: Request, user=Depends(get_verified_user 'name': connection.get('name', ''), } for connection in connections - if connection.get('enabled', True) and has_connection_access(user, connection, user_group_ids) + if connection.get('enabled', True) and await has_connection_access(user, connection, user_group_ids) ] @@ -82,8 +82,8 @@ async def proxy_terminal( if connection is None: return JSONResponse({'error': f"Terminal server '{server_id}' not found"}, status_code=404) - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)} - if not has_connection_access(user, connection, user_group_ids): + user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)} + if not await has_connection_access(user, connection, user_group_ids): return JSONResponse({'error': 'Access denied'}, status_code=403) base_url = (connection.get('url') or '').rstrip('/') @@ -208,7 +208,7 @@ async def _resolve_authenticated_connection(ws: WebSocket, server_id: str): if data is None or 'id' not in data: await ws.close(code=4001, reason='Invalid token') return None - user = Users.get_user_by_id(data['id']) + user = await Users.get_user_by_id(data['id']) if user is None: await ws.close(code=4001, reason='User not found') return None @@ -227,8 +227,8 @@ async def _resolve_authenticated_connection(ws: WebSocket, server_id: str): await ws.close(code=4004, reason='Terminal server not found') return None - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)} - if not has_connection_access(user, connection, user_group_ids): + user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)} + if not await has_connection_access(user, connection, user_group_ids): await ws.close(code=4003, reason='Access denied') return None diff --git a/backend/open_webui/routers/tools.py b/backend/open_webui/routers/tools.py index 195a4eec3e..c61ef7752f 100644 --- a/backend/open_webui/routers/tools.py +++ b/backend/open_webui/routers/tools.py @@ -8,8 +8,8 @@ from open_webui.env import AIOHTTP_CLIENT_TIMEOUT from open_webui.models.groups import Groups from pydantic import BaseModel, HttpUrl from fastapi import APIRouter, Depends, HTTPException, Request, status -from sqlalchemy.orm import Session -from open_webui.internal.db import get_session +from sqlalchemy.ext.asyncio import AsyncSession +from open_webui.internal.db import get_async_session from open_webui.models.oauth_sessions import OAuthSessions @@ -46,11 +46,11 @@ log = logging.getLogger(__name__) router = APIRouter() -def get_tool_module(request, tool_id, load_from_db=True): +async def get_tool_module(request, tool_id, load_from_db=True): """ Get the tool module by its ID. """ - tool_module, _ = get_tool_module_from_cache(request, tool_id, load_from_db) + tool_module, _ = await get_tool_module_from_cache(request, tool_id, load_from_db) return tool_module @@ -65,12 +65,12 @@ def get_tool_module(request, tool_id, load_from_db=True): async def get_tools( request: Request, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): tools = [] # Local Tools - for tool in Tools.get_tools(defer_content=True, db=db): + for tool in await Tools.get_tools(defer_content=True, db=db): tool_module = request.app.state.TOOLS.get(tool.id) if hasattr(request.app.state, 'TOOLS') else None tools.append( ToolUserResponse( @@ -159,7 +159,7 @@ async def get_tools( # Admin can see all tools return tools else: - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id, db=db)} + user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)} tools = [ tool for tool in tools @@ -173,7 +173,7 @@ async def get_tools( db=db, ) if str(tool.id).startswith('server:') - else AccessGrants.has_access( + else await AccessGrants.has_access( user_id=user.id, resource_type='tool', resource_id=tool.id, @@ -192,13 +192,13 @@ async def get_tools( @router.get('/list', response_model=list[ToolAccessResponse]) -async def get_tool_list(user=Depends(get_verified_user), db: Session = Depends(get_session)): +async def get_tool_list(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): if user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL: - tools = Tools.get_tools(defer_content=True, db=db) + tools = await Tools.get_tools(defer_content=True, db=db) else: - tools = Tools.get_tools_by_user_id(user.id, 'read', defer_content=True, db=db) + tools = await Tools.get_tools_by_user_id(user.id, 'read', defer_content=True, db=db) - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id, db=db)} + user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)} result = [] for tool in tools: @@ -298,9 +298,9 @@ async def load_tool_from_url(request: Request, form_data: LoadUrlForm, user=Depe async def export_tools( request: Request, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - if user.role != 'admin' and not has_permission( + if user.role != 'admin' and not await has_permission( user.id, 'workspace.tools_export', request.app.state.config.USER_PERMISSIONS, @@ -312,9 +312,9 @@ async def export_tools( ) if user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL: - return Tools.get_tools(db=db) + return await Tools.get_tools(db=db) else: - return Tools.get_tools_by_user_id(user.id, 'read', db=db) + return await Tools.get_tools_by_user_id(user.id, 'read', db=db) ############################ @@ -327,11 +327,11 @@ async def create_new_tools( request: Request, form_data: ToolForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): if user.role != 'admin' and not ( - has_permission(user.id, 'workspace.tools', request.app.state.config.USER_PERMISSIONS, db=db) - or has_permission( + await has_permission(user.id, 'workspace.tools', request.app.state.config.USER_PERMISSIONS, db=db) + or await has_permission( user.id, 'workspace.tools_import', request.app.state.config.USER_PERMISSIONS, @@ -351,10 +351,10 @@ async def create_new_tools( form_data.id = form_data.id.lower() - tools = Tools.get_tool_by_id(form_data.id, db=db) + tools = await Tools.get_tool_by_id(form_data.id, db=db) if tools is None: try: - form_data.access_grants = filter_allowed_access_grants( + form_data.access_grants = await filter_allowed_access_grants( request.app.state.config.USER_PERMISSIONS, user.id, user.role, @@ -363,14 +363,14 @@ async def create_new_tools( ) form_data.content = replace_imports(form_data.content) - tool_module, frontmatter = load_tool_module_by_id(form_data.id, content=form_data.content) + tool_module, frontmatter = await load_tool_module_by_id(form_data.id, content=form_data.content) form_data.meta.manifest = frontmatter TOOLS = request.app.state.TOOLS TOOLS[form_data.id] = tool_module specs = get_tool_specs(TOOLS[form_data.id]) - tools = Tools.insert_new_tool(user.id, form_data, specs, db=db) + tools = await Tools.insert_new_tool(user.id, form_data, specs, db=db) tool_cache_dir = CACHE_DIR / 'tools' / form_data.id tool_cache_dir.mkdir(parents=True, exist_ok=True) @@ -401,14 +401,14 @@ async def create_new_tools( @router.get('/id/{id}', response_model=Optional[ToolAccessResponse]) -async def get_tools_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - tools = Tools.get_tool_by_id(id, db=db) +async def get_tools_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + tools = await Tools.get_tool_by_id(id, db=db) if tools: if ( user.role == 'admin' or tools.user_id == user.id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='tool', resource_id=tools.id, @@ -421,7 +421,7 @@ async def get_tools_by_id(id: str, user=Depends(get_verified_user), db: Session write_access=( (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL) or user.id == tools.user_id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='tool', resource_id=tools.id, @@ -453,9 +453,9 @@ async def update_tools_by_id( id: str, form_data: ToolForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - tools = Tools.get_tool_by_id(id, db=db) + tools = await Tools.get_tool_by_id(id, db=db) if not tools: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, @@ -465,7 +465,7 @@ async def update_tools_by_id( # Is the user the original creator, in a group with write access, or an admin if ( tools.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='tool', resource_id=tools.id, @@ -481,7 +481,7 @@ async def update_tools_by_id( try: form_data.content = replace_imports(form_data.content) - tool_module, frontmatter = load_tool_module_by_id(id, content=form_data.content) + tool_module, frontmatter = await load_tool_module_by_id(id, content=form_data.content) form_data.meta.manifest = frontmatter TOOLS = request.app.state.TOOLS @@ -489,7 +489,7 @@ async def update_tools_by_id( specs = get_tool_specs(TOOLS[id]) - form_data.access_grants = filter_allowed_access_grants( + form_data.access_grants = await filter_allowed_access_grants( request.app.state.config.USER_PERMISSIONS, user.id, user.role, @@ -503,7 +503,7 @@ async def update_tools_by_id( } log.debug(updated) - tools = Tools.update_tool_by_id(id, updated, db=db) + tools = await Tools.update_tool_by_id(id, updated, db=db) if tools: return tools @@ -535,9 +535,9 @@ async def update_tool_access_by_id( id: str, form_data: ToolAccessGrantsForm, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - tools = Tools.get_tool_by_id(id, db=db) + tools = await Tools.get_tool_by_id(id, db=db) if not tools: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -546,7 +546,7 @@ async def update_tool_access_by_id( if ( tools.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='tool', resource_id=tools.id, @@ -560,7 +560,7 @@ async def update_tool_access_by_id( detail=ERROR_MESSAGES.UNAUTHORIZED, ) - form_data.access_grants = filter_allowed_access_grants( + form_data.access_grants = await filter_allowed_access_grants( request.app.state.config.USER_PERMISSIONS, user.id, user.role, @@ -568,9 +568,9 @@ async def update_tool_access_by_id( 'sharing.public_tools', ) - AccessGrants.set_access_grants('tool', id, form_data.access_grants, db=db) + await AccessGrants.set_access_grants('tool', id, form_data.access_grants, db=db) - return Tools.get_tool_by_id(id, db=db) + return await Tools.get_tool_by_id(id, db=db) ############################ @@ -583,9 +583,9 @@ async def delete_tools_by_id( request: Request, id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - tools = Tools.get_tool_by_id(id, db=db) + tools = await Tools.get_tool_by_id(id, db=db) if not tools: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, @@ -594,7 +594,7 @@ async def delete_tools_by_id( if ( tools.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='tool', resource_id=tools.id, @@ -608,7 +608,7 @@ async def delete_tools_by_id( detail=ERROR_MESSAGES.UNAUTHORIZED, ) - result = Tools.delete_tool_by_id(id, db=db) + result = await Tools.delete_tool_by_id(id, db=db) if result: TOOLS = request.app.state.TOOLS if id in TOOLS: @@ -623,8 +623,8 @@ async def delete_tools_by_id( @router.get('/id/{id}/valves', response_model=Optional[dict]) -async def get_tools_valves_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - tools = Tools.get_tool_by_id(id, db=db) +async def get_tools_valves_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + tools = await Tools.get_tool_by_id(id, db=db) if not tools: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -633,7 +633,7 @@ async def get_tools_valves_by_id(id: str, user=Depends(get_verified_user), db: S if ( tools.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='tool', resource_id=tools.id, @@ -648,7 +648,7 @@ async def get_tools_valves_by_id(id: str, user=Depends(get_verified_user), db: S ) try: - valves = Tools.get_tool_valves_by_id(id, db=db) + valves = await Tools.get_tool_valves_by_id(id, db=db) return valves except Exception as e: raise HTTPException( @@ -667,9 +667,9 @@ async def get_tools_valves_spec_by_id( request: Request, id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - tools = Tools.get_tool_by_id(id, db=db) + tools = await Tools.get_tool_by_id(id, db=db) if not tools: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -678,7 +678,7 @@ async def get_tools_valves_spec_by_id( if ( tools.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='tool', resource_id=tools.id, @@ -695,7 +695,7 @@ async def get_tools_valves_spec_by_id( if id in request.app.state.TOOLS: tools_module = request.app.state.TOOLS[id] else: - tools_module, _ = load_tool_module_by_id(id) + tools_module, _ = await load_tool_module_by_id(id) request.app.state.TOOLS[id] = tools_module if hasattr(tools_module, 'Valves'): @@ -718,9 +718,9 @@ async def update_tools_valves_by_id( id: str, form_data: dict, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - tools = Tools.get_tool_by_id(id, db=db) + tools = await Tools.get_tool_by_id(id, db=db) if not tools: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, @@ -729,7 +729,7 @@ async def update_tools_valves_by_id( if ( tools.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='tool', resource_id=tools.id, @@ -746,7 +746,7 @@ async def update_tools_valves_by_id( if id in request.app.state.TOOLS: tools_module = request.app.state.TOOLS[id] else: - tools_module, _ = load_tool_module_by_id(id) + tools_module, _ = await load_tool_module_by_id(id) request.app.state.TOOLS[id] = tools_module if not hasattr(tools_module, 'Valves'): @@ -760,7 +760,7 @@ async def update_tools_valves_by_id( form_data = {k: v for k, v in form_data.items() if v is not None} valves = Valves(**form_data) valves_dict = valves.model_dump(exclude_unset=True) - Tools.update_tool_valves_by_id(id, valves_dict, db=db) + await Tools.update_tool_valves_by_id(id, valves_dict, db=db) return valves_dict except Exception as e: log.exception(f'Failed to update tool valves by id {id}: {e}') @@ -776,8 +776,8 @@ async def update_tools_valves_by_id( @router.get('/id/{id}/valves/user', response_model=Optional[dict]) -async def get_tools_user_valves_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - tools = Tools.get_tool_by_id(id, db=db) +async def get_tools_user_valves_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + tools = await Tools.get_tool_by_id(id, db=db) if not tools: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -786,7 +786,7 @@ async def get_tools_user_valves_by_id(id: str, user=Depends(get_verified_user), if ( tools.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='tool', resource_id=tools.id, @@ -801,7 +801,7 @@ async def get_tools_user_valves_by_id(id: str, user=Depends(get_verified_user), ) try: - user_valves = Tools.get_user_valves_by_id_and_user_id(id, user.id, db=db) + user_valves = await Tools.get_user_valves_by_id_and_user_id(id, user.id, db=db) return user_valves except Exception as e: raise HTTPException( @@ -815,9 +815,9 @@ async def get_tools_user_valves_spec_by_id( request: Request, id: str, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - tools = Tools.get_tool_by_id(id, db=db) + tools = await Tools.get_tool_by_id(id, db=db) if not tools: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -826,7 +826,7 @@ async def get_tools_user_valves_spec_by_id( if ( tools.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='tool', resource_id=tools.id, @@ -843,7 +843,7 @@ async def get_tools_user_valves_spec_by_id( if id in request.app.state.TOOLS: tools_module = request.app.state.TOOLS[id] else: - tools_module, _ = load_tool_module_by_id(id) + tools_module, _ = await load_tool_module_by_id(id) request.app.state.TOOLS[id] = tools_module if hasattr(tools_module, 'UserValves'): @@ -861,9 +861,9 @@ async def update_tools_user_valves_by_id( id: str, form_data: dict, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - tools = Tools.get_tool_by_id(id, db=db) + tools = await Tools.get_tool_by_id(id, db=db) if not tools: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -872,7 +872,7 @@ async def update_tools_user_valves_by_id( if ( tools.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='tool', resource_id=tools.id, @@ -889,7 +889,7 @@ async def update_tools_user_valves_by_id( if id in request.app.state.TOOLS: tools_module = request.app.state.TOOLS[id] else: - tools_module, _ = load_tool_module_by_id(id) + tools_module, _ = await load_tool_module_by_id(id) request.app.state.TOOLS[id] = tools_module if hasattr(tools_module, 'UserValves'): @@ -899,7 +899,7 @@ async def update_tools_user_valves_by_id( form_data = {k: v for k, v in form_data.items() if v is not None} user_valves = UserValves(**form_data) user_valves_dict = user_valves.model_dump(exclude_unset=True) - Tools.update_user_valves_by_id_and_user_id(id, user.id, user_valves_dict, db=db) + await Tools.update_user_valves_by_id_and_user_id(id, user.id, user_valves_dict, db=db) return user_valves_dict except Exception as e: log.exception(f'Failed to update user valves by id {id}: {e}') diff --git a/backend/open_webui/routers/users.py b/backend/open_webui/routers/users.py index 0ccc20185e..091143a5d5 100644 --- a/backend/open_webui/routers/users.py +++ b/backend/open_webui/routers/users.py @@ -1,6 +1,6 @@ import logging from typing import Optional -from sqlalchemy.orm import Session +from sqlalchemy.ext.asyncio import AsyncSession import base64 import io @@ -30,7 +30,7 @@ from open_webui.models.users import ( from open_webui.constants import ERROR_MESSAGES from open_webui.env import STATIC_DIR -from open_webui.internal.db import get_session +from open_webui.internal.db import get_async_session from open_webui.utils.auth import ( @@ -63,7 +63,7 @@ async def get_users( direction: Optional[str] = None, page: Optional[int] = 1, user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): limit = PAGE_ITEM_COUNT @@ -80,14 +80,14 @@ async def get_users( filter['direction'] = direction - result = Users.get_users(filter=filter, skip=skip, limit=limit, db=db) + result = await Users.get_users(filter=filter, skip=skip, limit=limit, db=db) users = result['users'] total = result['total'] # Fetch groups for all users in a single query to avoid N+1 user_ids = [user.id for user in users] - user_groups = Groups.get_groups_by_member_ids(user_ids, db=db) + user_groups = await Groups.get_groups_by_member_ids(user_ids, db=db) return { 'users': [ @@ -106,9 +106,9 @@ async def get_users( @router.get('/all', response_model=UserInfoListResponse) async def get_all_users( user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - return Users.get_users(db=db) + return await Users.get_users(db=db) @router.get('/search', response_model=UserInfoListResponse) @@ -118,7 +118,7 @@ async def search_users( direction: Optional[str] = None, page: Optional[int] = 1, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): limit = PAGE_ITEM_COUNT @@ -133,7 +133,7 @@ async def search_users( if direction: filter['direction'] = direction - return Users.get_users(filter=filter, skip=skip, limit=limit, db=db) + return await Users.get_users(filter=filter, skip=skip, limit=limit, db=db) ############################ @@ -142,8 +142,8 @@ async def search_users( @router.get('/groups') -async def get_user_groups(user=Depends(get_verified_user), db: Session = Depends(get_session)): - return Groups.get_groups_by_member_id(user.id, db=db) +async def get_user_groups(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + return await Groups.get_groups_by_member_id(user.id, db=db) ############################ @@ -155,9 +155,9 @@ async def get_user_groups(user=Depends(get_verified_user), db: Session = Depends async def get_user_permissisions( request: Request, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): - user_permissions = get_permissions(user.id, request.app.state.config.USER_PERMISSIONS, db=db) + user_permissions = await get_permissions(user.id, request.app.state.config.USER_PERMISSIONS, db=db) return user_permissions @@ -272,8 +272,8 @@ async def update_default_user_permissions(request: Request, form_data: UserPermi @router.get('/user/settings', response_model=Optional[UserSettings]) -async def get_user_settings_by_session_user(user=Depends(get_verified_user), db: Session = Depends(get_session)): - user = Users.get_user_by_id(user.id, db=db) +async def get_user_settings_by_session_user(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + user = await Users.get_user_by_id(user.id, db=db) if user: return user.settings else: @@ -293,7 +293,7 @@ async def update_user_settings_by_session_user( request: Request, form_data: UserSettings, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): updated_user_settings = form_data.model_dump() ui_settings = updated_user_settings.get('ui') @@ -301,7 +301,7 @@ async def update_user_settings_by_session_user( user.role != 'admin' and ui_settings is not None and 'toolServers' in ui_settings.keys() - and not has_permission( + and not await has_permission( user.id, 'features.direct_tool_servers', request.app.state.config.USER_PERMISSIONS, @@ -310,7 +310,7 @@ async def update_user_settings_by_session_user( # If the user is not an admin and does not have permission to use tool servers, remove the key updated_user_settings['ui'].pop('toolServers', None) - user = Users.update_user_settings_by_id(user.id, updated_user_settings, db=db) + user = await Users.update_user_settings_by_id(user.id, updated_user_settings, db=db) if user: return user.settings else: @@ -329,14 +329,14 @@ async def update_user_settings_by_session_user( async def get_user_status_by_session_user( request: Request, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): if not request.app.state.config.ENABLE_USER_STATUS: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.ACTION_PROHIBITED, ) - user = Users.get_user_by_id(user.id, db=db) + user = await Users.get_user_by_id(user.id, db=db) if user: return user else: @@ -356,16 +356,16 @@ async def update_user_status_by_session_user( request: Request, form_data: UserStatus, user=Depends(get_verified_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): if not request.app.state.config.ENABLE_USER_STATUS: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.ACTION_PROHIBITED, ) - user = Users.get_user_by_id(user.id, db=db) + user = await Users.get_user_by_id(user.id, db=db) if user: - user = Users.update_user_status_by_id(user.id, form_data, db=db) + user = await Users.update_user_status_by_id(user.id, form_data, db=db) return user else: raise HTTPException( @@ -380,8 +380,8 @@ async def update_user_status_by_session_user( @router.get('/user/info', response_model=Optional[dict]) -async def get_user_info_by_session_user(user=Depends(get_verified_user), db: Session = Depends(get_session)): - user = Users.get_user_by_id(user.id, db=db) +async def get_user_info_by_session_user(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + user = await Users.get_user_by_id(user.id, db=db) if user: return user.info else: @@ -398,14 +398,14 @@ async def get_user_info_by_session_user(user=Depends(get_verified_user), db: Ses @router.post('/user/info/update', response_model=Optional[dict]) async def update_user_info_by_session_user( - form_data: dict, user=Depends(get_verified_user), db: Session = Depends(get_session) + form_data: dict, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session) ): - user = Users.get_user_by_id(user.id, db=db) + user = await Users.get_user_by_id(user.id, db=db) if user: if user.info is None: user.info = {} - user = Users.update_user_by_id(user.id, {'info': {**user.info, **form_data}}, db=db) + user = await Users.update_user_by_id(user.id, {'info': {**user.info, **form_data}}, db=db) if user: return user.info else: @@ -435,12 +435,12 @@ class UserActiveResponse(UserStatus): @router.get('/{user_id}', response_model=UserActiveResponse) -async def get_user_by_id(user_id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)): +async def get_user_by_id(user_id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): # Check if user_id is a shared chat # If it is, get the user_id from the chat if user_id.startswith('shared-'): chat_id = user_id.replace('shared-', '') - chat = Chats.get_chat_by_id(chat_id) + chat = await Chats.get_chat_by_id(chat_id) if chat: user_id = chat.user_id else: @@ -449,14 +449,14 @@ async def get_user_by_id(user_id: str, user=Depends(get_admin_user), db: Session detail=ERROR_MESSAGES.USER_NOT_FOUND, ) - user = Users.get_user_by_id(user_id, db=db) + user = await Users.get_user_by_id(user_id, db=db) if user: - groups = Groups.get_groups_by_member_id(user_id, db=db) + groups = await Groups.get_groups_by_member_id(user_id, db=db) return UserActiveResponse( **{ **user.model_dump(), 'groups': [{'id': group.id, 'name': group.name} for group in groups], - 'is_active': Users.is_user_active(user_id, db=db), + 'is_active': await Users.is_user_active(user_id, db=db), } ) else: @@ -467,15 +467,15 @@ async def get_user_by_id(user_id: str, user=Depends(get_admin_user), db: Session @router.get('/{user_id}/info', response_model=UserInfoResponse) -async def get_user_info_by_id(user_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)): - user = Users.get_user_by_id(user_id, db=db) +async def get_user_info_by_id(user_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): + user = await Users.get_user_by_id(user_id, db=db) if user: - groups = Groups.get_groups_by_member_id(user_id, db=db) + groups = await Groups.get_groups_by_member_id(user_id, db=db) return UserInfoResponse( **{ **user.model_dump(), 'groups': [{'id': group.id, 'name': group.name} for group in groups], - 'is_active': Users.is_user_active(user_id, db=db), + 'is_active': await Users.is_user_active(user_id, db=db), } ) else: @@ -486,8 +486,8 @@ async def get_user_info_by_id(user_id: str, user=Depends(get_verified_user), db: @router.get('/{user_id}/oauth/sessions') -async def get_user_oauth_sessions_by_id(user_id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)): - sessions = OAuthSessions.get_sessions_by_user_id(user_id, db=db) +async def get_user_oauth_sessions_by_id(user_id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): + sessions = await OAuthSessions.get_sessions_by_user_id(user_id, db=db) if sessions and len(sessions) > 0: return sessions else: @@ -503,8 +503,8 @@ async def get_user_oauth_sessions_by_id(user_id: str, user=Depends(get_admin_use @router.get('/{user_id}/profile/image') -def get_user_profile_image_by_id(user_id: str, user=Depends(get_verified_user)): - user = Users.get_user_by_id(user_id) +async def get_user_profile_image_by_id(user_id: str, user=Depends(get_verified_user)): + user = await Users.get_user_by_id(user_id) if user: if user.profile_image_url: # check if it's url or base64 @@ -542,10 +542,10 @@ def get_user_profile_image_by_id(user_id: str, user=Depends(get_verified_user)): @router.get('/{user_id}/active', response_model=dict) async def get_user_active_status_by_id( - user_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session) + user_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session) ): return { - 'active': Users.is_user_active(user_id, db=db), + 'active': await Users.is_user_active(user_id, db=db), } @@ -559,11 +559,11 @@ async def update_user_by_id( user_id: str, form_data: UserUpdateForm, session_user=Depends(get_admin_user), - db: Session = Depends(get_session), + db: AsyncSession = Depends(get_async_session), ): # Prevent modification of the primary admin user by other admins try: - first_user = Users.get_first_user(db=db) + first_user = await Users.get_first_user(db=db) if first_user: if user_id == first_user.id: if session_user.id != user_id: @@ -587,11 +587,11 @@ async def update_user_by_id( detail='Could not verify primary admin status.', ) - user = Users.get_user_by_id(user_id, db=db) + user = await Users.get_user_by_id(user_id, db=db) if user: if form_data.email.lower() != user.email: - email_user = Users.get_user_by_email(form_data.email.lower(), db=db) + email_user = await Users.get_user_by_email(form_data.email.lower(), db=db) if email_user: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -605,10 +605,10 @@ async def update_user_by_id( raise HTTPException(400, detail=str(e)) hashed = get_password_hash(form_data.password) - Auths.update_user_password_by_id(user_id, hashed, db=db) + await Auths.update_user_password_by_id(user_id, hashed, db=db) - Auths.update_email_by_id(user_id, form_data.email.lower(), db=db) - updated_user = Users.update_user_by_id( + await Auths.update_email_by_id(user_id, form_data.email.lower(), db=db) + updated_user = await Users.update_user_by_id( user_id, { 'role': form_data.role, @@ -639,10 +639,10 @@ async def update_user_by_id( @router.delete('/{user_id}', response_model=bool) -async def delete_user_by_id(user_id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)): +async def delete_user_by_id(user_id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): # Prevent deletion of the primary admin user try: - first_user = Users.get_first_user(db=db) + first_user = await Users.get_first_user(db=db) if first_user and user_id == first_user.id: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, @@ -656,7 +656,7 @@ async def delete_user_by_id(user_id: str, user=Depends(get_admin_user), db: Sess ) if user.id != user_id: - result = Auths.delete_auth_by_id(user_id, db=db) + result = await Auths.delete_auth_by_id(user_id, db=db) if result: return True @@ -679,5 +679,5 @@ async def delete_user_by_id(user_id: str, user=Depends(get_admin_user), db: Sess @router.get('/{user_id}/groups') -async def get_user_groups_by_id(user_id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)): - return Groups.get_groups_by_member_id(user_id, db=db) +async def get_user_groups_by_id(user_id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): + return await Groups.get_groups_by_member_id(user_id, db=db) diff --git a/backend/open_webui/socket/main.py b/backend/open_webui/socket/main.py index 80e8b5be1c..d2ddd90b19 100644 --- a/backend/open_webui/socket/main.py +++ b/backend/open_webui/socket/main.py @@ -333,7 +333,7 @@ async def connect(sid, environ, auth): data = decode_token(auth['token']) if data is not None and 'id' in data: - user = Users.get_user_by_id(data['id']) + user = await Users.get_user_by_id(data['id']) if user: SESSION_POOL[sid] = { @@ -361,7 +361,7 @@ async def user_join(sid, data): if data is None or 'id' not in data: return - user = Users.get_user_by_id(data['id']) + user = await Users.get_user_by_id(data['id']) if not user: return @@ -381,8 +381,8 @@ 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 has_permission(user.id, 'features.channels'): - channels = Channels.get_channels_by_user_id(user.id) + if user.role == 'admin' or await has_permission(user.id, 'features.channels'): + channels = await Channels.get_channels_by_user_id(user.id) log.debug(f'{channels=}') for channel in channels: await sio.enter_room(sid, f'channel:{channel.id}') @@ -395,7 +395,7 @@ async def heartbeat(sid, data): user = SESSION_POOL.get(sid) if user: SESSION_POOL[sid] = {**user, 'last_seen_at': int(time.time())} - await asyncio.to_thread(Users.update_last_active_by_id, user['id']) + await Users.update_last_active_by_id(user['id']) @sio.on('join-channels') @@ -408,13 +408,13 @@ async def join_channel(sid, data): if data is None or 'id' not in data: return - user = Users.get_user_by_id(data['id']) + user = await Users.get_user_by_id(data['id']) if not user: return # Join all the channels only if user has channels permission - if user.role == 'admin' or has_permission(user.id, 'features.channels'): - channels = Channels.get_channels_by_user_id(user.id) + if user.role == 'admin' or await has_permission(user.id, 'features.channels'): + channels = await Channels.get_channels_by_user_id(user.id) log.debug(f'{channels=}') for channel in channels: await sio.enter_room(sid, f'channel:{channel.id}') @@ -430,11 +430,11 @@ async def join_note(sid, data): if token_data is None or 'id' not in token_data: return - user = Users.get_user_by_id(token_data['id']) + user = await Users.get_user_by_id(token_data['id']) if not user: return - note = Notes.get_note_by_id(data['note_id']) + 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}') return @@ -442,7 +442,7 @@ async def join_note(sid, data): if ( user.role != 'admin' and user.id != note.user_id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='note', resource_id=note.id, @@ -488,7 +488,7 @@ async def channel_events(sid, data): room=room, ) elif event_type == 'last_read_at': - Channels.update_member_last_read_at(data['channel_id'], user['id']) + await Channels.update_member_last_read_at(data['channel_id'], user['id']) @sio.on('events:chat') @@ -501,7 +501,7 @@ async def chat_events(sid, data): event_type = event_data.get('type') if event_type == 'last_read_at': - await asyncio.to_thread(Chats.update_chat_last_read_at_by_id, data['chat_id'], user['id']) + await Chats.update_chat_last_read_at_by_id(data['chat_id'], user['id']) def normalize_document_id(document_id: str) -> str: @@ -529,7 +529,7 @@ async def ydoc_document_join(sid, data): if document_id.startswith('note:'): note_id = document_id.split(':')[1] - note = Notes.get_note_by_id(note_id) + note = await Notes.get_note_by_id(note_id) if not note: log.error(f'Note {note_id} not found') return @@ -537,7 +537,7 @@ async def ydoc_document_join(sid, data): if ( user.get('role') != 'admin' and user.get('id') != note.user_id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.get('id'), resource_type='note', resource_id=note.id, @@ -602,7 +602,7 @@ async def document_save_handler(document_id, data, user): if document_id.startswith('note:'): note_id = document_id.split(':')[1] - note = Notes.get_note_by_id(note_id) + note = await Notes.get_note_by_id(note_id) if not note: log.error(f'Note {note_id} not found') return @@ -610,7 +610,7 @@ async def document_save_handler(document_id, data, user): if ( user.get('role') != 'admin' and user.get('id') != note.user_id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.get('id'), resource_type='note', resource_id=note.id, @@ -620,7 +620,7 @@ async def document_save_handler(document_id, data, user): log.error(f'User {user.get("id")} does not have write access to note {note_id}') return - Notes.update_note_by_id(note_id, NoteUpdateForm(data=data)) + await Notes.update_note_by_id(note_id, NoteUpdateForm(data=data)) @sio.on('ydoc:document:state') @@ -793,7 +793,7 @@ async def disconnect(sid): # print(f"Unknown session ID {sid} disconnected") -def get_event_emitter(request_info, update_db=True): +async def get_event_emitter(request_info, update_db=True): async def __event_emitter__(event_data): user_id = request_info['user_id'] chat_id = request_info['chat_id'] @@ -813,16 +813,14 @@ def get_event_emitter(request_info, update_db=True): event_type = event_data.get('type') if event_type == 'status': - await asyncio.to_thread( - Chats.add_message_status_to_chat_by_id_and_message_id, + await Chats.add_message_status_to_chat_by_id_and_message_id( request_info['chat_id'], request_info['message_id'], event_data.get('data', {}), ) elif event_type == 'message': - message = await asyncio.to_thread( - Chats.get_message_by_id_and_message_id, + message = await Chats.get_message_by_id_and_message_id( request_info['chat_id'], request_info['message_id'], ) @@ -831,8 +829,7 @@ def get_event_emitter(request_info, update_db=True): content = message.get('content', '') content += event_data.get('data', {}).get('content', '') - await asyncio.to_thread( - Chats.upsert_message_to_chat_by_id_and_message_id, + await Chats.upsert_message_to_chat_by_id_and_message_id( request_info['chat_id'], request_info['message_id'], { @@ -843,8 +840,7 @@ def get_event_emitter(request_info, update_db=True): elif event_type == 'replace': content = event_data.get('data', {}).get('content', '') - await asyncio.to_thread( - Chats.upsert_message_to_chat_by_id_and_message_id, + await Chats.upsert_message_to_chat_by_id_and_message_id( request_info['chat_id'], request_info['message_id'], { @@ -853,8 +849,7 @@ def get_event_emitter(request_info, update_db=True): ) elif event_type == 'embeds': - message = await asyncio.to_thread( - Chats.get_message_by_id_and_message_id, + message = await Chats.get_message_by_id_and_message_id( request_info['chat_id'], request_info['message_id'], ) @@ -862,8 +857,7 @@ def get_event_emitter(request_info, update_db=True): embeds = event_data.get('data', {}).get('embeds', []) embeds.extend(message.get('embeds', [])) - await asyncio.to_thread( - Chats.upsert_message_to_chat_by_id_and_message_id, + await Chats.upsert_message_to_chat_by_id_and_message_id( request_info['chat_id'], request_info['message_id'], { @@ -872,8 +866,7 @@ def get_event_emitter(request_info, update_db=True): ) elif event_type == 'files': - message = await asyncio.to_thread( - Chats.get_message_by_id_and_message_id, + message = await Chats.get_message_by_id_and_message_id( request_info['chat_id'], request_info['message_id'], ) @@ -881,8 +874,7 @@ def get_event_emitter(request_info, update_db=True): files = event_data.get('data', {}).get('files', []) files.extend(message.get('files', [])) - await asyncio.to_thread( - Chats.upsert_message_to_chat_by_id_and_message_id, + await Chats.upsert_message_to_chat_by_id_and_message_id( request_info['chat_id'], request_info['message_id'], { @@ -893,8 +885,7 @@ def get_event_emitter(request_info, update_db=True): elif event_type in ('source', 'citation'): data = event_data.get('data', {}) if data.get('type') is None: - message = await asyncio.to_thread( - Chats.get_message_by_id_and_message_id, + message = await Chats.get_message_by_id_and_message_id( request_info['chat_id'], request_info['message_id'], ) @@ -902,8 +893,7 @@ def get_event_emitter(request_info, update_db=True): sources = message.get('sources', []) sources.append(data) - await asyncio.to_thread( - Chats.upsert_message_to_chat_by_id_and_message_id, + await Chats.upsert_message_to_chat_by_id_and_message_id( request_info['chat_id'], request_info['message_id'], { @@ -917,7 +907,7 @@ def get_event_emitter(request_info, update_db=True): return None -def get_event_call(request_info): +async def get_event_call(request_info): async def __event_caller__(event_data): response = await sio.call( 'events', diff --git a/backend/open_webui/tools/builtin.py b/backend/open_webui/tools/builtin.py index d6032eb900..58af934372 100644 --- a/backend/open_webui/tools/builtin.py +++ b/backend/open_webui/tools/builtin.py @@ -250,7 +250,7 @@ async def generate_image( # Persist files to DB if chat context is available if __chat_id__ and __message_id__ and images: - db_files = Chats.add_message_files_by_id_and_message_id( + db_files = await Chats.add_message_files_by_id_and_message_id( __chat_id__, __message_id__, image_files, @@ -317,7 +317,7 @@ async def edit_image( # Persist files to DB if chat context is available if __chat_id__ and __message_id__ and images: - db_files = Chats.add_message_files_by_id_and_message_id( + db_files = await Chats.add_message_files_by_id_and_message_id( __chat_id__, __message_id__, image_files, @@ -473,14 +473,14 @@ async def execute_code( from open_webui.models.users import Users from open_webui.utils.files import get_image_url_from_base64 - user = Users.get_user_by_id(__user__['id']) + user = await Users.get_user_by_id(__user__['id']) # Extract and upload images from stdout if stdout and isinstance(stdout, str): stdout_lines = stdout.split('\n') for idx, line in enumerate(stdout_lines): if 'data:image/png;base64' in line: - image_url = get_image_url_from_base64( + image_url = await get_image_url_from_base64( __request__, line, __metadata__ or {}, @@ -495,7 +495,7 @@ async def execute_code( result_lines = result.split('\n') for idx, line in enumerate(result_lines): if 'data:image/png;base64' in line: - image_url = get_image_url_from_base64( + image_url = await get_image_url_from_base64( __request__, line, __metadata__ or {}, @@ -650,7 +650,7 @@ async def delete_memory( try: user = UserModel(**__user__) if __user__ else None - result = Memories.delete_memory_by_id_and_user_id(memory_id, user.id) + result = await Memories.delete_memory_by_id_and_user_id(memory_id, user.id) if result: VECTOR_DB_CLIENT.delete(collection_name=f'user-memory-{user.id}', ids=[memory_id]) @@ -680,7 +680,7 @@ async def list_memories( try: user = UserModel(**__user__) if __user__ else None - memories = Memories.get_memories_by_user_id(user.id) + memories = await Memories.get_memories_by_user_id(user.id) if memories: result = [ @@ -730,9 +730,9 @@ async def search_notes( try: user_id = __user__.get('id') - user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id)] + user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)] - result = Notes.search_notes( + result = await Notes.search_notes( user_id=user_id, filter={ 'query': query, @@ -808,18 +808,18 @@ async def view_note( return json.dumps({'error': 'User context not available'}) try: - note = Notes.get_note_by_id(note_id) + note = await Notes.get_note_by_id(note_id) if not note: return json.dumps({'error': 'Note not found'}) # Check access permission user_id = __user__.get('id') - user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id)] + user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)] from open_webui.models.access_grants import AccessGrants - if note.user_id != user_id and not AccessGrants.has_access( + if note.user_id != user_id and not await AccessGrants.has_access( user_id=user_id, resource_type='note', resource_id=note.id, @@ -878,7 +878,7 @@ async def write_note( access_grants=[], # Private by default - only owner can access ) - new_note = Notes.insert_new_note(user_id, form) + new_note = await Notes.insert_new_note(user_id, form) if not new_note: return json.dumps({'error': 'Failed to create note'}) @@ -921,18 +921,18 @@ async def replace_note_content( try: from open_webui.models.notes import NoteUpdateForm - note = Notes.get_note_by_id(note_id) + note = await Notes.get_note_by_id(note_id) if not note: return json.dumps({'error': 'Note not found'}) # Check write permission user_id = __user__.get('id') - user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id)] + user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)] from open_webui.models.access_grants import AccessGrants - if note.user_id != user_id and not AccessGrants.has_access( + if note.user_id != user_id and not await AccessGrants.has_access( user_id=user_id, resource_type='note', resource_id=note.id, @@ -947,7 +947,7 @@ async def replace_note_content( update_data['title'] = title form = NoteUpdateForm(**update_data) - updated_note = Notes.update_note_by_id(note_id, form) + updated_note = await Notes.update_note_by_id(note_id, form) if not updated_note: return json.dumps({'error': 'Failed to update note'}) @@ -998,7 +998,7 @@ async def search_chats( try: user_id = __user__.get('id') - chats = Chats.get_chats_by_user_id_and_search_text( + chats = await Chats.get_chats_by_user_id_and_search_text( user_id=user_id, search_text=query, include_archived=False, @@ -1073,7 +1073,7 @@ async def view_chat( try: user_id = __user__.get('id') - chat = Chats.get_chat_by_id_and_user_id(chat_id, user_id) + chat = await Chats.get_chat_by_id_and_user_id(chat_id, user_id) if not chat: return json.dumps({'error': 'Chat not found or access denied'}) @@ -1145,7 +1145,7 @@ async def search_channels( user_id = __user__.get('id') # Get all channels the user has access to - all_channels = Channels.get_channels_by_user_id(user_id) + all_channels = await Channels.get_channels_by_user_id(user_id) # Filter by query lower_query = query.lower() @@ -1201,7 +1201,7 @@ async def search_channel_messages( user_id = __user__.get('id') # Get all channels the user has access to - user_channels = Channels.get_channels_by_user_id(user_id) + user_channels = await Channels.get_channels_by_user_id(user_id) channel_ids = [c.id for c in user_channels] channel_map = {c.id: c for c in user_channels} @@ -1280,12 +1280,12 @@ async def view_channel_message( return json.dumps({'error': 'Message not found'}) # Verify user has access to the channel - channel = Channels.get_channel_by_id(message.channel_id) + channel = await Channels.get_channel_by_id(message.channel_id) if not channel: return json.dumps({'error': 'Channel not found'}) # Check if user has access to the channel - user_channels = Channels.get_channels_by_user_id(user_id) + user_channels = await Channels.get_channels_by_user_id(user_id) channel_ids = [c.id for c in user_channels] if message.channel_id not in channel_ids: @@ -1342,11 +1342,11 @@ async def view_channel_thread( return json.dumps({'error': 'Message not found'}) # Verify user has access to the channel - channel = Channels.get_channel_by_id(parent_message.channel_id) + channel = await Channels.get_channel_by_id(parent_message.channel_id) if not channel: return json.dumps({'error': 'Channel not found'}) - user_channels = Channels.get_channels_by_user_id(user_id) + user_channels = await Channels.get_channels_by_user_id(user_id) channel_ids = [c.id for c in user_channels] if parent_message.channel_id not in channel_ids: @@ -1427,9 +1427,9 @@ async def list_knowledge_bases( from open_webui.models.knowledge import Knowledges user_id = __user__.get('id') - user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id)] + user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)] - result = Knowledges.search_knowledge_bases( + result = await Knowledges.search_knowledge_bases( user_id, filter={ 'query': '', @@ -1442,7 +1442,7 @@ async def list_knowledge_bases( knowledge_bases = [] for knowledge_base in result.items: - files = Knowledges.get_files_by_id(knowledge_base.id) + files = await Knowledges.get_files_by_id(knowledge_base.id) file_count = len(files) if files else 0 knowledge_bases.append( @@ -1486,9 +1486,9 @@ async def search_knowledge_bases( from open_webui.models.knowledge import Knowledges user_id = __user__.get('id') - user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id)] + user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)] - result = Knowledges.search_knowledge_bases( + result = await Knowledges.search_knowledge_bases( user_id, filter={ 'query': query, @@ -1501,7 +1501,7 @@ async def search_knowledge_bases( knowledge_bases = [] for knowledge_base in result.items: - files = Knowledges.get_files_by_id(knowledge_base.id) + files = await Knowledges.get_files_by_id(knowledge_base.id) file_count = len(files) if files else 0 knowledge_bases.append( @@ -1552,7 +1552,7 @@ async def search_knowledge_files( user_id = __user__.get('id') user_role = __user__.get('role', 'user') - user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id)] + user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)] # When model has attached knowledge, scope to attached KBs/files only if __model_knowledge__: @@ -1577,14 +1577,14 @@ async def search_knowledge_files( # Search within attached KBs for kb_id in attached_kb_ids: - knowledge = Knowledges.get_knowledge_by_id(kb_id) + knowledge = await Knowledges.get_knowledge_by_id(kb_id) if not knowledge: continue if not ( user_role == 'admin' or knowledge.user_id == user_id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user_id, resource_type='knowledge', resource_id=knowledge.id, @@ -1594,7 +1594,7 @@ async def search_knowledge_files( ): continue - result = Knowledges.search_files_by_id( + result = await Knowledges.search_files_by_id( knowledge_id=kb_id, user_id=user_id, filter={'query': query}, @@ -1617,7 +1617,7 @@ async def search_knowledge_files( if not knowledge_id and attached_file_ids: query_lower = query.lower() if query else '' for file_id in attached_file_ids: - file = Files.get_file_by_id(file_id) + file = await Files.get_file_by_id(file_id) if file and (not query_lower or query_lower in file.filename.lower()): all_files.append( { @@ -1633,7 +1633,7 @@ async def search_knowledge_files( # No attached knowledge - search all accessible KBs if knowledge_id: - result = Knowledges.search_files_by_id( + result = await Knowledges.search_files_by_id( knowledge_id=knowledge_id, user_id=user_id, filter={'query': query}, @@ -1641,7 +1641,7 @@ async def search_knowledge_files( limit=count, ) else: - result = Knowledges.search_knowledge_files( + result = await Knowledges.search_knowledge_files( filter={ 'query': query, 'user_id': user_id, @@ -1719,7 +1719,7 @@ async def view_file( user_id = __user__.get('id') user_role = __user__.get('role', 'user') - file = Files.get_file_by_id(file_id) + file = await Files.get_file_by_id(file_id) if not file: return json.dumps({'error': 'File not found'}) @@ -1729,7 +1729,7 @@ async def view_file( and not any( item.get('type') == 'file' and item.get('id') == file_id for item in (__model_knowledge__ or []) ) - and not has_access_to_file( + and not await has_access_to_file( file_id=file_id, access_type='read', user=UserModel(**__user__), @@ -1811,14 +1811,14 @@ async def view_knowledge_file( user_id = __user__.get('id') user_role = __user__.get('role', 'user') - user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id)] + user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)] - file = Files.get_file_by_id(file_id) + file = await Files.get_file_by_id(file_id) if not file: return json.dumps({'error': 'File not found'}) # Check access via any KB containing this file - knowledges = Knowledges.get_knowledges_by_file_id(file_id) + knowledges = await Knowledges.get_knowledges_by_file_id(file_id) has_knowledge_access = False knowledge_info = None @@ -1826,7 +1826,7 @@ async def view_knowledge_file( if ( user_role == 'admin' or knowledge_base.user_id == user_id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user_id, resource_type='knowledge', resource_id=knowledge_base.id, @@ -1903,7 +1903,7 @@ async def list_knowledge( user_id = __user__.get('id') user_role = __user__.get('role', 'user') - user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id)] + user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)] knowledge_bases = [] files = [] @@ -1914,11 +1914,11 @@ async def list_knowledge( item_id = item.get('id') if item_type == 'collection': - knowledge = Knowledges.get_knowledge_by_id(item_id) + knowledge = await Knowledges.get_knowledge_by_id(item_id) if knowledge and ( user_role == 'admin' or knowledge.user_id == user_id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user_id, resource_type='knowledge', resource_id=knowledge.id, @@ -1926,7 +1926,7 @@ async def list_knowledge( user_group_ids=set(user_group_ids), ) ): - kb_files = Knowledges.get_files_by_id(knowledge.id) + kb_files = await Knowledges.get_files_by_id(knowledge.id) file_count = len(kb_files) if kb_files else 0 kb_entry = { @@ -1943,7 +1943,7 @@ async def list_knowledge( knowledge_bases.append(kb_entry) elif item_type == 'file': - file = Files.get_file_by_id(item_id) + file = await Files.get_file_by_id(item_id) if file: files.append( { @@ -1954,11 +1954,11 @@ async def list_knowledge( ) elif item_type == 'note': - note = Notes.get_note_by_id(item_id) + note = await Notes.get_note_by_id(item_id) if note and ( user_role == 'admin' or note.user_id == user_id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user_id, resource_type='note', resource_id=note.id, @@ -2036,7 +2036,7 @@ async def query_knowledge_files( user_id = __user__.get('id') user_role = __user__.get('role', 'user') - user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id)] + user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)] embedding_function = __request__.app.state.EMBEDDING_FUNCTION if not embedding_function: @@ -2053,11 +2053,11 @@ async def query_knowledge_files( if item_type == 'collection': # Knowledge base - use KB ID as collection name - knowledge = Knowledges.get_knowledge_by_id(item_id) + knowledge = await Knowledges.get_knowledge_by_id(item_id) if knowledge and ( user_role == 'admin' or knowledge.user_id == user_id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user_id, resource_type='knowledge', resource_id=knowledge.id, @@ -2069,17 +2069,17 @@ async def query_knowledge_files( elif item_type == 'file': # Individual file - use file-{id} as collection name - file = Files.get_file_by_id(item_id) + file = await Files.get_file_by_id(item_id) if file: collection_names.append(f'file-{item_id}') elif item_type == 'note': # Note - always return full content as context - note = Notes.get_note_by_id(item_id) + note = await Notes.get_note_by_id(item_id) if note and ( user_role == 'admin' or note.user_id == user_id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user_id, resource_type='note', resource_id=note.id, @@ -2099,11 +2099,11 @@ async def query_knowledge_files( elif knowledge_ids: # User specified specific KBs for knowledge_id in knowledge_ids: - knowledge = Knowledges.get_knowledge_by_id(knowledge_id) + knowledge = await Knowledges.get_knowledge_by_id(knowledge_id) if knowledge and ( user_role == 'admin' or knowledge.user_id == user_id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user_id, resource_type='knowledge', resource_id=knowledge.id, @@ -2114,7 +2114,7 @@ async def query_knowledge_files( collection_names.append(knowledge_id) else: # No model knowledge and no specific IDs - search all accessible KBs - result = Knowledges.search_knowledge_bases( + result = await Knowledges.search_knowledge_bases( user_id, filter={ 'query': '', @@ -2193,7 +2193,7 @@ async def query_knowledge_bases( from open_webui.retrieval.vector.factory import VECTOR_DB_CLIENT user_id = __user__.get('id') - user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id)] + user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)] query_embedding = await __request__.app.state.EMBEDDING_FUNCTION(query) # Min-heap of (distance, knowledge_base_id) - only holds top `count` results @@ -2203,7 +2203,7 @@ async def query_knowledge_bases( page_size = 100 while True: - accessible_knowledge_bases = Knowledges.search_knowledge_bases( + accessible_knowledge_bases = await Knowledges.search_knowledge_bases( user_id, filter={'user_id': user_id, 'group_ids': user_group_ids}, skip=page_offset, @@ -2247,7 +2247,7 @@ async def query_knowledge_bases( matching_knowledge_bases = [] for distance, knowledge_base_id in sorted_results: - knowledge_base = Knowledges.get_knowledge_by_id(knowledge_base_id) + knowledge_base = await Knowledges.get_knowledge_by_id(knowledge_base_id) if knowledge_base: matching_knowledge_bases.append( { @@ -2295,7 +2295,7 @@ async def view_skill( user_id = __user__.get('id') # Direct DB lookup by id (case-insensitive since IDs are stored lowercase) - skill = Skills.get_skill_by_id(id.lower()) + skill = await Skills.get_skill_by_id(id.lower()) if not skill or not skill.is_active: return json.dumps({'error': f"Skill '{id}' not found"}) @@ -2303,8 +2303,8 @@ async def view_skill( # Check user access user_role = __user__.get('role', 'user') if user_role != 'admin' and skill.user_id != user_id: - user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id)] - if not AccessGrants.has_access( + user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)] + if not await AccessGrants.has_access( user_id=user_id, resource_type='skill', resource_id=skill.id, @@ -2393,7 +2393,7 @@ async def tasks( if tasks is None: # Read-only - return current list - all_tasks = Chats.get_chat_tasks_by_id(__chat_id__) + all_tasks = await Chats.get_chat_tasks_by_id(__chat_id__) elif overwrite: # Full replacement - validate and write all_tasks = [] @@ -2417,7 +2417,7 @@ async def tasks( ) else: # Partial update - merge by id - existing_tasks = Chats.get_chat_tasks_by_id(__chat_id__) + existing_tasks = await Chats.get_chat_tasks_by_id(__chat_id__) existing_by_id = {t['id']: t for t in existing_tasks} seen_ids = set() @@ -2460,7 +2460,7 @@ async def tasks( # Persist to DB and emit (skip for read-only) if tasks is not None: - Chats.update_chat_tasks_by_id(__chat_id__, all_tasks) + await Chats.update_chat_tasks_by_id(__chat_id__, all_tasks) if __event_emitter__: await __event_emitter__( @@ -2542,7 +2542,7 @@ async def create_automation( from open_webui.utils.automations import validate_rrule, next_run_ns, next_n_runs_ns user_id = __user__.get('id') - user = Users.get_user_by_id(user_id) + user = await Users.get_user_by_id(user_id) if not user: return json.dumps({'error': 'User not found'}) @@ -2569,7 +2569,7 @@ async def create_automation( is_active=True, ) - automation = Automations.insert(user_id, form, next_run_ns(rrule, tz=tz)) + automation = await Automations.insert(user_id, form, next_run_ns(rrule, tz=tz)) return json.dumps( { @@ -2618,9 +2618,9 @@ async def update_automation( from open_webui.utils.automations import validate_rrule, next_run_ns, next_n_runs_ns user_id = __user__.get('id') - user = Users.get_user_by_id(user_id) + user = await Users.get_user_by_id(user_id) - automation = Automations.get_by_id(automation_id) + automation = await Automations.get_by_id(automation_id) if not automation: return json.dumps({'error': 'Automation not found'}) if automation.user_id != user_id: @@ -2650,7 +2650,7 @@ async def update_automation( is_active=automation.is_active, ) - updated = Automations.update(automation_id, form, next_run_ns(new_rrule, tz=tz)) + updated = await Automations.update(automation_id, form, next_run_ns(new_rrule, tz=tz)) return json.dumps( { @@ -2693,9 +2693,9 @@ async def list_automations( from open_webui.utils.automations import next_n_runs_ns user_id = __user__.get('id') - user = Users.get_user_by_id(user_id) + user = await Users.get_user_by_id(user_id) - result = Automations.search_automations( + result = await Automations.search_automations( user_id=user_id, status=status, skip=0, @@ -2753,16 +2753,16 @@ async def toggle_automation( from open_webui.utils.automations import next_run_ns user_id = __user__.get('id') - user = Users.get_user_by_id(user_id) + user = await Users.get_user_by_id(user_id) - automation = Automations.get_by_id(automation_id) + automation = await Automations.get_by_id(automation_id) if not automation: return json.dumps({'error': 'Automation not found'}) if automation.user_id != user_id: return json.dumps({'error': 'Access denied'}) rrule = automation.data.get('rrule', '') - toggled = Automations.toggle( + toggled = await Automations.toggle( automation_id, next_run_ns(rrule, tz=user.timezone if user else None), ) @@ -2803,15 +2803,15 @@ async def delete_automation( user_id = __user__.get('id') - automation = Automations.get_by_id(automation_id) + automation = await Automations.get_by_id(automation_id) if not automation: return json.dumps({'error': 'Automation not found'}) if automation.user_id != user_id: return json.dumps({'error': 'Access denied'}) name = automation.name - AutomationRuns.delete_by_automation(automation_id) - Automations.delete(automation_id) + await AutomationRuns.delete_by_automation(automation_id) + await Automations.delete(automation_id) return json.dumps( { diff --git a/backend/open_webui/utils/access_control/__init__.py b/backend/open_webui/utils/access_control/__init__.py index 9c91371384..41c888d441 100644 --- a/backend/open_webui/utils/access_control/__init__.py +++ b/backend/open_webui/utils/access_control/__init__.py @@ -11,7 +11,7 @@ from open_webui.models.access_grants import ( ) from open_webui.config import DEFAULT_USER_PERMISSIONS -from sqlalchemy.orm import Session +from sqlalchemy.ext.asyncio import AsyncSession def fill_missing_permissions(permissions: dict[str, Any], default_permissions: dict[str, Any]) -> dict[str, Any]: @@ -28,10 +28,10 @@ def fill_missing_permissions(permissions: dict[str, Any], default_permissions: d return permissions -def get_permissions( +async def get_permissions( user_id: str, default_permissions: dict[str, Any], - db: Session | None = None, + db: AsyncSession | None = None, ) -> dict[str, Any]: """ Get all permissions for a user by combining the permissions of all groups the user is a member of. @@ -53,7 +53,7 @@ def get_permissions( permissions[key] = permissions[key] or value # Use the most permissive value (True > False) return permissions - user_groups = Groups.get_groups_by_member_id(user_id, db=db) + user_groups = await Groups.get_groups_by_member_id(user_id, db=db) # Deep copy default permissions to avoid modifying the original dict permissions = json.loads(json.dumps(default_permissions)) @@ -68,11 +68,11 @@ def get_permissions( return permissions -def has_permission( +async def has_permission( user_id: str, permission_key: str, default_permissions: dict[str, Any] = {}, - db: Session | None = None, + db: AsyncSession | None = None, ) -> bool: """ Check if a user has a specific permission by checking the group permissions @@ -93,7 +93,7 @@ def has_permission( permission_hierarchy = permission_key.split('.') # Retrieve user group permissions - user_groups = Groups.get_groups_by_member_id(user_id, db=db) + user_groups = await Groups.get_groups_by_member_id(user_id, db=db) for group in user_groups: if get_permission(group.permissions or {}, permission_hierarchy): @@ -104,12 +104,12 @@ def has_permission( return get_permission(default_permissions, permission_hierarchy) -def has_access( +async def has_access( user_id: str, permission: str = 'read', access_grants: list | None = None, user_group_ids: set[str] | None = None, - db: Session | None = None, + db: AsyncSession | None = None, ) -> bool: """ Check if a user has the specified permission using an in-memory access_grants list. @@ -126,7 +126,7 @@ def has_access( return False if user_group_ids is None: - user_groups = Groups.get_groups_by_member_id(user_id, db=db) + user_groups = await Groups.get_groups_by_member_id(user_id, db=db) user_group_ids = {group.id for group in user_groups} for grant in access_grants: @@ -144,7 +144,7 @@ def has_access( return False -def has_connection_access( +async def has_connection_access( user: UserModel, connection: dict, user_group_ids: set[str] | None = None, @@ -163,10 +163,10 @@ def has_connection_access( return True if user_group_ids is None: - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)} + user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)} access_grants = (connection.get('config') or {}).get('access_grants', []) - return has_access(user.id, 'read', access_grants, user_group_ids) + return await has_access(user.id, 'read', access_grants, user_group_ids) def migrate_access_control(data: dict, ac_key: str = 'access_control', grants_key: str = 'access_grants') -> None: @@ -210,13 +210,13 @@ def migrate_access_control(data: dict, ac_key: str = 'access_control', grants_ke data.pop(ac_key, None) -def filter_allowed_access_grants( +async def filter_allowed_access_grants( default_permissions: dict[str, Any], user_id: str, user_role: str, access_grants: list, public_permission_key: str, - db: Session | None = None, + db: AsyncSession | None = None, ) -> list: """ Checks if the user has the required permissions to grant access to a resource. @@ -228,7 +228,7 @@ def filter_allowed_access_grants( # Check if user can share publicly if ( has_public_read_access_grant(access_grants) or has_public_write_access_grant(access_grants) - ) and not has_permission( + ) and not await has_permission( user_id, public_permission_key, default_permissions, @@ -246,7 +246,7 @@ def filter_allowed_access_grants( ] # Strip individual user sharing if user lacks permission - if has_user_access_grant(access_grants) and not has_permission( + if has_user_access_grant(access_grants) and not await has_permission( user_id, 'access_grants.allow_users', default_permissions, @@ -257,7 +257,7 @@ def filter_allowed_access_grants( return access_grants -def check_model_access( +async def check_model_access( user: UserModel, model_info, bypass_filter: bool = False, @@ -270,7 +270,7 @@ def check_model_access( Args: user: The authenticated user. - model_info: The model record from Models.get_model_by_id(), + model_info: The model record from await Models.get_model_by_id(), or None if the model is not registered. bypass_filter: If True, skip all access checks (used by internal callers and BYPASS_MODEL_ACCESS_CONTROL). @@ -284,10 +284,10 @@ def check_model_access( if user.role == 'user': from open_webui.models.access_grants import AccessGrants - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)} + user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)} if not ( user.id == model_info.user_id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='model', resource_id=model_info.id, diff --git a/backend/open_webui/utils/access_control/files.py b/backend/open_webui/utils/access_control/files.py index a7e35fd506..5e7efb5b26 100644 --- a/backend/open_webui/utils/access_control/files.py +++ b/backend/open_webui/utils/access_control/files.py @@ -9,16 +9,16 @@ from open_webui.models.groups import Groups from open_webui.models.models import Models from open_webui.models.access_grants import AccessGrants -from sqlalchemy.orm import Session +from sqlalchemy.ext.asyncio import AsyncSession log = logging.getLogger(__name__) -def has_access_to_file( +async def has_access_to_file( file_id: str | None, access_type: str, user: UserModel, - db: Session | None = None, + db: AsyncSession | None = None, ) -> bool: """ Check if a user has the specified access to a file through any of: @@ -30,7 +30,7 @@ def has_access_to_file( NOTE: This does NOT check direct file ownership — callers should check file.user_id == user.id separately before calling this. """ - file = Files.get_file_by_id(file_id, db=db) + file = await Files.get_file_by_id(file_id, db=db) log.debug(f'Checking if user has {access_type} access to file') if not file: return False @@ -40,10 +40,10 @@ def has_access_to_file( return True # Check if the file is associated with any knowledge bases the user has access to - knowledge_bases = Knowledges.get_knowledges_by_file_id(file_id, db=db) - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id, db=db)} + knowledge_bases = await Knowledges.get_knowledges_by_file_id(file_id, db=db) + user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)} for knowledge_base in knowledge_bases: - if knowledge_base.user_id == user.id or AccessGrants.has_access( + if knowledge_base.user_id == user.id or await AccessGrants.has_access( user_id=user.id, resource_type='knowledge', resource_id=knowledge_base.id, @@ -55,24 +55,24 @@ def has_access_to_file( knowledge_base_id = file.meta.get('collection_name') if file.meta else None if knowledge_base_id: - knowledge_bases = Knowledges.get_knowledge_bases_by_user_id(user.id, access_type, db=db) + knowledge_bases = await Knowledges.get_knowledge_bases_by_user_id(user.id, access_type, db=db) for knowledge_base in knowledge_bases: if knowledge_base.id == knowledge_base_id: return True # Check if the file is associated with any channels the user has access to - channels = Channels.get_channels_by_file_id_and_user_id(file_id, user.id, db=db) + channels = await Channels.get_channels_by_file_id_and_user_id(file_id, user.id, db=db) if access_type == 'read' and channels: return True # Check if the file is associated with any chats the user has access to # TODO: Granular access control for chats - chats = Chats.get_shared_chats_by_file_id(file_id, db=db) + chats = await Chats.get_shared_chats_by_file_id(file_id, db=db) if chats: return True # Check if the file is directly attached to a shared workspace model - for model in Models.get_models_by_user_id(user.id, permission=access_type, db=db): + for model in await Models.get_models_by_user_id(user.id, permission=access_type, db=db): knowledge_items = getattr(model.meta, 'knowledge', None) or [] for item in knowledge_items: if isinstance(item, dict) and item.get('type') == 'file' and item.get('id') == file.id: diff --git a/backend/open_webui/utils/actions.py b/backend/open_webui/utils/actions.py index 5c5712fa0f..7b1789580b 100644 --- a/backend/open_webui/utils/actions.py +++ b/backend/open_webui/utils/actions.py @@ -26,7 +26,7 @@ async def chat_action(request: Request, action_id: str, form_data: dict, user: A else: sub_action_id = None - action = Functions.get_function_by_id(action_id) + action = await Functions.get_function_by_id(action_id) if not action: raise Exception(f'Action not found: {action_id}') @@ -47,7 +47,7 @@ async def chat_action(request: Request, action_id: str, form_data: dict, user: A raise Exception('Model not found') model = models[model_id] - __event_emitter__ = get_event_emitter( + __event_emitter__ = await get_event_emitter( { 'chat_id': data['chat_id'], 'message_id': data['id'], @@ -55,7 +55,7 @@ async def chat_action(request: Request, action_id: str, form_data: dict, user: A 'user_id': user.id, } ) - __event_call__ = get_event_call( + __event_call__ = await get_event_call( { 'chat_id': data['chat_id'], 'message_id': data['id'], @@ -64,10 +64,10 @@ async def chat_action(request: Request, action_id: str, form_data: dict, user: A } ) - function_module, _, _ = get_function_module_from_cache(request, action_id) + function_module, _, _ = await get_function_module_from_cache(request, action_id) if hasattr(function_module, 'valves') and hasattr(function_module, 'Valves'): - valves = Functions.get_function_valves_by_id(action_id) + valves = await Functions.get_function_valves_by_id(action_id) function_module.valves = function_module.Valves(**(valves if valves else {})) if hasattr(function_module, 'action'): @@ -98,7 +98,7 @@ async def chat_action(request: Request, action_id: str, form_data: dict, user: A try: if hasattr(function_module, 'UserValves'): __user__['valves'] = function_module.UserValves( - **Functions.get_user_valves_by_id_and_user_id(action_id, user.id) + **await Functions.get_user_valves_by_id_and_user_id(action_id, user.id) ) except Exception as e: log.exception(f'Failed to get user values: {e}') @@ -111,7 +111,7 @@ async def chat_action(request: Request, action_id: str, form_data: dict, user: A data = action(**params) # Process action result for Rich UI embeds (HTMLResponse, tuple with headers) - processed_result, _, action_embeds = process_tool_result( + processed_result, _, action_embeds = await process_tool_result( request, action_id, data, diff --git a/backend/open_webui/utils/auth.py b/backend/open_webui/utils/auth.py index fcdadf9acf..32e7db3423 100644 --- a/backend/open_webui/utils/auth.py +++ b/backend/open_webui/utils/auth.py @@ -324,7 +324,7 @@ async def get_current_user( # auth by api key if token.startswith('sk-'): - user = get_current_user_by_api_key(request, token) + user = await get_current_user_by_api_key(request, token) # Add user info to current span if ENABLE_OTEL: @@ -356,7 +356,7 @@ async def get_current_user( detail='Invalid token', ) - user = Users.get_user_by_id(data['id']) + user = await Users.get_user_by_id(data['id']) if user is None: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, @@ -382,10 +382,10 @@ async def get_current_user( current_span.set_attribute('client.user.role', user.role) current_span.set_attribute('client.auth.type', 'jwt') - # Refresh the user's last active timestamp asynchronously - # to prevent blocking the request - if background_tasks: - background_tasks.add_task(Users.update_last_active_by_id, user.id) + # Refresh the user's last active timestamp + # Fire-and-forget via asyncio.create_task to avoid blocking + import asyncio + asyncio.create_task(Users.update_last_active_by_id(user.id)) return user else: raise HTTPException( @@ -407,9 +407,9 @@ async def get_current_user( raise e -def get_current_user_by_api_key(request, api_key: str): +async def get_current_user_by_api_key(request, api_key: str): # Each function call manages its own short-lived session internally - user = Users.get_user_by_api_key(api_key) + user = await Users.get_user_by_api_key(api_key) if user is None: raise HTTPException( @@ -419,7 +419,7 @@ def get_current_user_by_api_key(request, api_key: str): if not request.state.enable_api_keys or ( user.role != 'admin' - and not has_permission( + and not await has_permission( user.id, 'features.api_keys', request.app.state.config.USER_PERMISSIONS, @@ -438,7 +438,7 @@ def get_current_user_by_api_key(request, api_key: str): current_span.set_attribute('client.user.role', user.role) current_span.set_attribute('client.auth.type', 'api_key') - Users.update_last_active_by_id(user.id) + await Users.update_last_active_by_id(user.id) return user @@ -460,7 +460,7 @@ def get_admin_user(user=Depends(get_current_user)): return user -def create_admin_user(email: str, password: str, name: str = 'Admin'): +async def create_admin_user(email: str, password: str, name: str = 'Admin'): """ Create an admin user from environment variables. Used for headless/automated deployments. @@ -470,14 +470,14 @@ def create_admin_user(email: str, password: str, name: str = 'Admin'): if not email or not password: return None - if Users.has_users(): + if await Users.has_users(): log.debug('Users already exist, skipping admin creation') return None log.info(f'Creating admin account from environment variables: {email}') try: hashed = get_password_hash(password) - user = Auths.insert_new_auth( + user = await Auths.insert_new_auth( email=email.lower(), password=hashed, name=name, diff --git a/backend/open_webui/utils/automations.py b/backend/open_webui/utils/automations.py index 5d28307f7a..262430dcf0 100644 --- a/backend/open_webui/utils/automations.py +++ b/backend/open_webui/utils/automations.py @@ -126,7 +126,7 @@ async def automation_worker_loop(app) -> None: while True: try: with get_db() as db: - batch = Automations.claim_due(int(time.time_ns()), limit=10, db=db) + batch = await Automations.claim_due(int(time.time_ns()), limit=10, db=db) if batch: log.info(f'Claimed {len(batch)} due automation(s)') for automation in batch: @@ -283,9 +283,9 @@ async def execute_automation(app, automation: AutomationModel) -> None: (filters, model params, knowledge/RAG, tools, DB saves, webhooks). """ try: - user = Users.get_user_by_id(automation.user_id) + user = await Users.get_user_by_id(automation.user_id) if not user: - _record_run(automation.id, 'error', error='User not found') + await _record_run(automation.id, 'error', error='User not found') return prompt = prompt_template(automation.data['prompt'], user) @@ -297,7 +297,7 @@ async def execute_automation(app, automation: AutomationModel) -> None: assistant_msg_id = str(uuid4()) # Create the chat with user message (same structure as frontend) - chat = Chats.insert_new_chat( + chat = await Chats.insert_new_chat( automation.user_id, ChatForm( chat={ @@ -336,7 +336,7 @@ async def execute_automation(app, automation: AutomationModel) -> None: ) if not chat: - _record_run(automation.id, 'error', error='Failed to create chat') + await _record_run(automation.id, 'error', error='Failed to create chat') return # Notify frontend to refresh chat list @@ -404,11 +404,11 @@ async def execute_automation(app, automation: AutomationModel) -> None: room=f'user:{automation.user_id}', ) - _record_run(automation.id, 'success', chat_id=chat.id) + await _record_run(automation.id, 'success', chat_id=chat.id) except Exception as e: log.exception(f'Automation {automation.id} failed') - _record_run(automation.id, 'error', error=str(e)[:4000]) + await _record_run(automation.id, 'error', error=str(e)[:4000]) #################### @@ -416,7 +416,7 @@ async def execute_automation(app, automation: AutomationModel) -> None: #################### -def _record_run( +async def _record_run( automation_id: str, status: str, chat_id: str = None, @@ -424,4 +424,4 @@ def _record_run( ): """Insert a run record into automation_run.""" with get_db() as db: - AutomationRuns.insert(automation_id, status, chat_id=chat_id, error=error, db=db) + await AutomationRuns.insert(automation_id, status, chat_id=chat_id, error=error, db=db) diff --git a/backend/open_webui/utils/chat.py b/backend/open_webui/utils/chat.py index 9a9e810331..3539d57c86 100644 --- a/backend/open_webui/utils/chat.py +++ b/backend/open_webui/utils/chat.py @@ -72,7 +72,7 @@ async def generate_direct_chat_completion( session_id = metadata.get('session_id') request_id = str(uuid.uuid4()) # Generate a unique request ID - event_caller = get_event_call(metadata) + event_caller = await get_event_call(metadata) channel = f'{user_id}:{session_id}:{request_id}' logging.info(f'WebSocket channel: {channel}') @@ -199,7 +199,7 @@ async def generate_chat_completion( # Check if user has access to the model if not bypass_filter and user.role == 'user': try: - check_model_access(user, model) + await check_model_access(user, model) except Exception as e: raise e @@ -343,8 +343,8 @@ async def chat_completed(request: Request, form_data: dict, user: Any): } extra_params = { - '__event_emitter__': get_event_emitter(metadata), - '__event_call__': get_event_call(metadata), + '__event_emitter__': await get_event_emitter(metadata), + '__event_call__': await get_event_call(metadata), '__user__': user.model_dump() if isinstance(user, UserModel) else {}, '__metadata__': metadata, '__request__': request, @@ -352,8 +352,8 @@ async def chat_completed(request: Request, form_data: dict, user: Any): } try: - filter_ids = get_sorted_filter_ids(request, model, metadata.get('filter_ids', [])) - filter_functions = Functions.get_functions_by_ids(filter_ids) + filter_ids = await get_sorted_filter_ids(request, model, metadata.get('filter_ids', [])) + filter_functions = await Functions.get_functions_by_ids(filter_ids) result, _ = await process_filter_functions( request=request, diff --git a/backend/open_webui/utils/embeddings.py b/backend/open_webui/utils/embeddings.py index 251b5edf7e..1717886326 100644 --- a/backend/open_webui/utils/embeddings.py +++ b/backend/open_webui/utils/embeddings.py @@ -68,7 +68,7 @@ async def generate_embeddings( # Access filtering if not getattr(request.state, 'direct', False): if not bypass_filter and user.role == 'user': - check_model_access(user, model) + await check_model_access(user, model) # Ollama backend — use /api/embed which supports batch input natively if model.get('owned_by') == 'ollama': diff --git a/backend/open_webui/utils/files.py b/backend/open_webui/utils/files.py index 06bec33250..ef7900a1ce 100644 --- a/backend/open_webui/utils/files.py +++ b/backend/open_webui/utils/files.py @@ -31,7 +31,7 @@ BASE64_IMAGE_URL_PREFIX = re.compile(r'data:image/\w+;base64,', re.IGNORECASE) MARKDOWN_IMAGE_URL_PATTERN = re.compile(r'!\[(.*?)\]\((.+?)\)', re.IGNORECASE) -def get_image_base64_from_url(url: str) -> Optional[str]: +async def get_image_base64_from_url(url: str) -> Optional[str]: try: if url.startswith('http'): # Validate URL to prevent SSRF attacks against local/private networks @@ -44,7 +44,7 @@ def get_image_base64_from_url(url: str) -> Optional[str]: content_type = response.headers.get('Content-Type', 'image/png') return f'data:{content_type};base64,{encoded_string}' else: - file = Files.get_file_by_id(url) + file = await Files.get_file_by_id(url) if not file: return None @@ -64,13 +64,13 @@ def get_image_base64_from_url(url: str) -> Optional[str]: return None -def get_image_url_from_base64(request, base64_image_string, metadata, user): +async def get_image_url_from_base64(request, base64_image_string, metadata, user): if BASE64_IMAGE_URL_PREFIX.match(base64_image_string): image_url = '' # Extract base64 image data from the line image_data, content_type = get_image_data(base64_image_string) if image_data is not None: - _, image_url = upload_image( + _, image_url = await upload_image( request, image_data, content_type, @@ -82,17 +82,26 @@ def get_image_url_from_base64(request, base64_image_string, metadata, user): return None -def convert_markdown_base64_images(request, content: str, metadata, user): - def replace(match): - base64_string = match.group(2) - MIN_REPLACEMENT_URL_LENGTH = 1024 - if len(base64_string) > MIN_REPLACEMENT_URL_LENGTH: - url = get_image_url_from_base64(request, base64_string, metadata, user) - if url: - return f'![{match.group(1)}]({url})' - return match.group(0) +async def convert_markdown_base64_images(request, content: str, metadata, user): + MIN_REPLACEMENT_URL_LENGTH = 1024 + result_parts = [] + last_end = 0 - return MARKDOWN_IMAGE_URL_PATTERN.sub(replace, content) + for match in MARKDOWN_IMAGE_URL_PATTERN.finditer(content): + result_parts.append(content[last_end:match.start()]) + base64_string = match.group(2) + if len(base64_string) > MIN_REPLACEMENT_URL_LENGTH: + url = await get_image_url_from_base64(request, base64_string, metadata, user) + if url: + result_parts.append(f'![{match.group(1)}]({url})') + else: + result_parts.append(match.group(0)) + else: + result_parts.append(match.group(0)) + last_end = match.end() + + result_parts.append(content[last_end:]) + return ''.join(result_parts) def load_b64_audio_data(b64_str): @@ -110,7 +119,7 @@ def load_b64_audio_data(b64_str): return None, None -def upload_audio(request, audio_data, content_type, metadata, user): +async def upload_audio(request, audio_data, content_type, metadata, user): audio_format = mimetypes.guess_extension(content_type) file = UploadFile( file=io.BytesIO(audio_data), @@ -119,7 +128,7 @@ def upload_audio(request, audio_data, content_type, metadata, user): 'content-type': content_type, }, ) - file_item = upload_file_handler( + file_item = await upload_file_handler( request, file=file, metadata=metadata, @@ -130,13 +139,13 @@ def upload_audio(request, audio_data, content_type, metadata, user): return url -def get_audio_url_from_base64(request, base64_audio_string, metadata, user): +async def get_audio_url_from_base64(request, base64_audio_string, metadata, user): if 'data:audio/wav;base64' in base64_audio_string: audio_url = '' # Extract base64 audio data from the line audio_data, content_type = load_b64_audio_data(base64_audio_string) if audio_data is not None: - audio_url = upload_audio( + audio_url = await upload_audio( request, audio_data, content_type, @@ -147,16 +156,16 @@ def get_audio_url_from_base64(request, base64_audio_string, metadata, user): return None -def get_file_url_from_base64(request, base64_file_string, metadata, user): +async def get_file_url_from_base64(request, base64_file_string, metadata, user): if BASE64_IMAGE_URL_PREFIX.match(base64_file_string): - return get_image_url_from_base64(request, base64_file_string, metadata, user) + return await get_image_url_from_base64(request, base64_file_string, metadata, user) elif 'data:audio/wav;base64' in base64_file_string: - return get_audio_url_from_base64(request, base64_file_string, metadata, user) + return await get_audio_url_from_base64(request, base64_file_string, metadata, user) return None -def get_image_base64_from_file_id(id: str) -> Optional[str]: - file = Files.get_file_by_id(id) +async def get_image_base64_from_file_id(id: str) -> Optional[str]: + file = await Files.get_file_by_id(id) if not file: return None diff --git a/backend/open_webui/utils/filter.py b/backend/open_webui/utils/filter.py index df07dea4a1..50b1583088 100644 --- a/backend/open_webui/utils/filter.py +++ b/backend/open_webui/utils/filter.py @@ -10,44 +10,53 @@ from open_webui.models.functions import Functions log = logging.getLogger(__name__) -def get_function_module(request, function_id, load_from_db=True): +async def get_function_module(request, function_id, load_from_db=True): """ Get the function module by its ID. """ - function_module, _, _ = get_function_module_from_cache(request, function_id, load_from_db) + function_module, _, _ = await get_function_module_from_cache(request, function_id, load_from_db) return function_module -def get_sorted_filter_ids(request, model: dict, enabled_filter_ids: list = None): - def get_priority(function_id): +async def get_sorted_filter_ids(request, model: dict, enabled_filter_ids: list = None): + async def get_priority(function_id): try: - function_module = get_function_module(request, function_id) + function_module = await get_function_module(request, function_id) if function_module and hasattr(function_module, 'Valves'): - valves_db = Functions.get_function_valves_by_id(function_id) + valves_db = await Functions.get_function_valves_by_id(function_id) valves = function_module.Valves(**(valves_db if valves_db else {})) return getattr(valves, 'priority', 0) except Exception: pass return 0 - filter_ids = [function.id for function in Functions.get_global_filter_functions()] + filter_ids = [function.id for function in await Functions.get_global_filter_functions()] if 'info' in model and 'meta' in model['info']: filter_ids.extend(model['info']['meta'].get('filterIds', [])) filter_ids = list(set(filter_ids)) - active_filter_ids = {function.id for function in Functions.get_functions_by_type('filter', active_only=True)} + active_filter_ids = {function.id for function in await Functions.get_functions_by_type('filter', active_only=True)} - def get_active_status(filter_id): - function_module = get_function_module(request, filter_id) + async def get_active_status(filter_id): + function_module = await get_function_module(request, filter_id) if getattr(function_module, 'toggle', None): return filter_id in (enabled_filter_ids or set()) return True - active_filter_ids = {filter_id for filter_id in active_filter_ids if get_active_status(filter_id)} + # Pre-compute active status for each filter (async functions can't be used in set comprehensions) + resolved_active = {} + for filter_id in active_filter_ids: + resolved_active[filter_id] = await get_active_status(filter_id) + active_filter_ids = {fid for fid, is_active in resolved_active.items() if is_active} filter_ids = [fid for fid in filter_ids if fid in active_filter_ids] - filter_ids.sort(key=lambda fid: (get_priority(fid), fid)) + + # Pre-compute priorities (async functions can't be used in sort keys) + priorities = {} + for fid in filter_ids: + priorities[fid] = await get_priority(fid) + filter_ids.sort(key=lambda fid: (priorities.get(fid, 0), fid)) return filter_ids @@ -63,7 +72,7 @@ async def process_filter_functions(request, filter_functions, filter_type, form_ if not filter: continue - function_module = get_function_module(request, filter_id, load_from_db=(filter_type != 'stream')) + function_module = await get_function_module(request, filter_id, load_from_db=(filter_type != 'stream')) # Prepare handler function handler = getattr(function_module, filter_type, None) if not handler: @@ -75,7 +84,7 @@ async def process_filter_functions(request, filter_functions, filter_type, form_ # Apply valves to the function if hasattr(function_module, 'valves') and hasattr(function_module, 'Valves'): - valves = Functions.get_function_valves_by_id(filter_id) + valves = await Functions.get_function_valves_by_id(filter_id) function_module.valves = function_module.Valves(**(valves if valves else {})) try: @@ -100,7 +109,7 @@ async def process_filter_functions(request, filter_functions, filter_type, form_ if hasattr(function_module, 'UserValves'): try: params['__user__']['valves'] = function_module.UserValves( - **Functions.get_user_valves_by_id_and_user_id(filter_id, params['__user__']['id']) + **await Functions.get_user_valves_by_id_and_user_id(filter_id, params['__user__']['id']) ) except Exception as e: log.exception(f'Failed to get user values: {e}') diff --git a/backend/open_webui/utils/groups.py b/backend/open_webui/utils/groups.py index 90c4593cec..50099b2ee7 100644 --- a/backend/open_webui/utils/groups.py +++ b/backend/open_webui/utils/groups.py @@ -4,7 +4,7 @@ from open_webui.models.groups import Groups log = logging.getLogger(__name__) -def apply_default_group_assignment( +async def apply_default_group_assignment( default_group_id: str, user_id: str, db=None, @@ -18,6 +18,6 @@ def apply_default_group_assignment( """ if default_group_id: try: - Groups.add_users_to_group(default_group_id, [user_id], db=db) + await Groups.add_users_to_group(default_group_id, [user_id], db=db) except Exception as e: log.error(f'Failed to add user {user_id} to default group {default_group_id}: {e}') diff --git a/backend/open_webui/utils/mcp/client.py b/backend/open_webui/utils/mcp/client.py index beb2f15079..effe4b1637 100644 --- a/backend/open_webui/utils/mcp/client.py +++ b/backend/open_webui/utils/mcp/client.py @@ -44,7 +44,7 @@ def create_httpx_client(headers=None, timeout=None, auth=None): return _build_httpx_client(headers=headers, timeout=timeout, auth=auth, verify=True) -def create_insecure_httpx_client(headers=None, timeout=None, auth=None): +async def create_insecure_httpx_client(headers=None, timeout=None, auth=None): return _build_httpx_client(headers=headers, timeout=timeout, auth=auth, verify=False) diff --git a/backend/open_webui/utils/middleware.py b/backend/open_webui/utils/middleware.py index 0dedd7f2f6..fb4912bef5 100644 --- a/backend/open_webui/utils/middleware.py +++ b/backend/open_webui/utils/middleware.py @@ -936,7 +936,7 @@ def apply_source_context_to_messages( ) -def process_tool_result( +async def process_tool_result( request, tool_function_name, tool_result, @@ -1075,7 +1075,7 @@ def process_tool_result( pass tool_response.append(text) elif item.get('type') in ['image', 'audio']: - file_url = get_file_url_from_base64( + file_url = await get_file_url_from_base64( request, f'data:{item.get("mimeType")};base64,{item.get("data", item.get("blob", ""))}', { @@ -1304,7 +1304,7 @@ async def chat_completion_tools_handler( except Exception as e: tool_result = str(e) - tool_result, tool_result_files, tool_result_embeds = process_tool_result( + tool_result, tool_result_files, tool_result_embeds = await process_tool_result( request, tool_function_name, tool_result, @@ -1602,7 +1602,7 @@ def get_images_from_messages(message_list): return images -def get_image_urls(delta_images, request, metadata, user) -> list[str]: +async def get_image_urls(delta_images, request, metadata, user) -> list[str]: if not isinstance(delta_images, list): return [] @@ -1616,21 +1616,21 @@ def get_image_urls(delta_images, request, metadata, user) -> list[str]: continue if url.startswith('data:image/png;base64'): - url = get_image_url_from_base64(request, url, metadata, user) + url = await get_image_url_from_base64(request, url, metadata, user) image_urls.append(url) return image_urls -def add_file_context(messages: list, chat_id: str, user) -> list: +async def add_file_context(messages: list, chat_id: str, user) -> list: """ Add file URLs to messages for native function calling. """ if not chat_id or chat_id.startswith('local:'): return messages - chat = Chats.get_chat_by_id_and_user_id(chat_id, user.id) + chat = await Chats.get_chat_by_id_and_user_id(chat_id, user.id) if not chat: return messages @@ -1686,7 +1686,7 @@ async def chat_image_generation_handler(request: Request, form_data: dict, extra if chat_id.startswith('local:'): message_list = form_data.get('messages', []) else: - chat = Chats.get_chat_by_id_and_user_id(chat_id, user.id) + chat = await Chats.get_chat_by_id_and_user_id(chat_id, user.id) await __event_emitter__( { 'type': 'status', @@ -2066,12 +2066,12 @@ async def convert_url_images_to_base64(form_data): return form_data -def load_messages_from_db(chat_id: str, message_id: str) -> Optional[list[dict]]: +async def load_messages_from_db(chat_id: str, message_id: str) -> Optional[list[dict]]: """ Load the message chain from DB up to message_id, keeping only LLM-relevant fields (role, content, output). """ - messages_map = Chats.get_messages_map_by_chat_id(chat_id) + messages_map = await Chats.get_messages_map_by_chat_id(chat_id) if not messages_map: return None @@ -2149,7 +2149,7 @@ async def process_chat_payload(request, form_data, user, metadata, model): parent_message_id = metadata.get('parent_message_id') if chat_id and parent_message_id and not chat_id.startswith('local:'): - db_messages = load_messages_from_db(chat_id, parent_message_id) + db_messages = await load_messages_from_db(chat_id, parent_message_id) if db_messages: system_message = get_system_message(form_data.get('messages', [])) form_data['messages'] = [system_message, *db_messages] if system_message else db_messages @@ -2192,8 +2192,8 @@ async def process_chat_payload(request, form_data, user, metadata, model): form_data = await convert_url_images_to_base64(form_data) - event_emitter = get_event_emitter(metadata) - event_caller = get_event_call(metadata) + event_emitter = await get_event_emitter(metadata) + event_caller = await get_event_call(metadata) extra_params = { '__event_emitter__': event_emitter, @@ -2231,14 +2231,14 @@ async def process_chat_payload(request, form_data, user, metadata, model): chat_id = metadata.get('chat_id', None) folder_id = None if chat_id and user: - folder_id = Chats.get_chat_folder_id(chat_id, user.id) + folder_id = await Chats.get_chat_folder_id(chat_id, user.id) # Fallback: use folder_id from metadata (temporary chats have no DB record) if not folder_id: folder_id = metadata.get('folder_id', None) if folder_id and user: - folder = Folders.get_folder_by_id_and_user_id(folder_id, user.id) + folder = await Folders.get_folder_by_id_and_user_id(folder_id, user.id) if folder and folder.data: if 'system_prompt' in folder.data: @@ -2305,8 +2305,8 @@ async def process_chat_payload(request, form_data, user, metadata, model): raise e try: - filter_ids = get_sorted_filter_ids(request, model, metadata.get('filter_ids', [])) - filter_functions = Functions.get_functions_by_ids(filter_ids) + filter_ids = await get_sorted_filter_ids(request, model, metadata.get('filter_ids', [])) + filter_functions = await Functions.get_functions_by_ids(filter_ids) form_data, flags = await process_filter_functions( request=request, @@ -2399,12 +2399,13 @@ async def process_chat_payload(request, form_data, user, metadata, model): if all_skill_ids: from open_webui.models.skills import Skills as SkillsModel - accessible_skill_ids = {s.id for s in SkillsModel.get_skills_by_user_id(user.id, 'read')} - available_skills = [ - s - for sid in all_skill_ids - if sid in accessible_skill_ids and (s := SkillsModel.get_skill_by_id(sid)) and s.is_active - ] + accessible_skill_ids = {s.id for s in await SkillsModel.get_skills_by_user_id(user.id, 'read')} + available_skills = [] + for sid in all_skill_ids: + if sid in accessible_skill_ids: + s = await SkillsModel.get_skill_by_id(sid) + if s and s.is_active: + available_skills.append(s) skill_descriptions = '' for skill in available_skills: @@ -2441,7 +2442,7 @@ async def process_chat_payload(request, form_data, user, metadata, model): # Get folder files folder_id = file_item.get('id', None) if folder_id: - folder = Folders.get_folder_by_id_and_user_id(folder_id, user.id) + folder = await Folders.get_folder_by_id_and_user_id(folder_id, user.id) if folder and folder.data and 'files' in folder.data: files = [f for f in files if f.get('id', None) != folder_id] files = [*files, *folder.data['files']] @@ -2495,7 +2496,7 @@ async def process_chat_payload(request, form_data, user, metadata, model): continue # Check access control for MCP server - if not has_connection_access(user, mcp_server_connection): + if not await has_connection_access(user, mcp_server_connection): log.warning(f'Access denied to MCP server {server_id} for user {user.id}') continue @@ -2556,7 +2557,7 @@ async def process_chat_payload(request, form_data, user, metadata, model): tool_specs = await mcp_clients[server_id].list_tool_specs() for tool_spec in tool_specs: - def make_tool_function(client, function_name): + async def make_tool_function(client, function_name): async def tool_function(**kwargs): return await client.call_tool( function_name, @@ -2570,7 +2571,7 @@ async def process_chat_payload(request, form_data, user, metadata, model): # Skip this function continue - tool_function = make_tool_function(mcp_clients[server_id], tool_spec['name']) + tool_function = await make_tool_function(mcp_clients[server_id], tool_spec['name']) mcp_tools_dict[f'{server_id}_{tool_spec["name"]}'] = { 'spec': { @@ -2664,8 +2665,8 @@ async def process_chat_payload(request, form_data, user, metadata, model): if metadata.get('params', {}).get('function_calling') == 'native' and builtin_tools_enabled: # Add file context to user messages chat_id = metadata.get('chat_id') - form_data['messages'] = add_file_context(form_data.get('messages', []), chat_id, user) - builtin_tools = get_builtin_tools( + form_data['messages'] = await add_file_context(form_data.get('messages', []), chat_id, user) + builtin_tools = await get_builtin_tools( request, { **extra_params, @@ -2755,7 +2756,7 @@ async def process_chat_payload(request, form_data, user, metadata, model): return form_data, metadata, events -def get_event_emitter_and_caller(metadata): +async def get_event_emitter_and_caller(metadata): event_emitter = None event_caller = None @@ -2763,18 +2764,18 @@ def get_event_emitter_and_caller(metadata): # It broadcasts to user:{user_id} room AND persists to DB, # so it works for backend-initiated calls (automations, API). if metadata.get('chat_id') and metadata.get('message_id'): - event_emitter = get_event_emitter(metadata) + event_emitter = await get_event_emitter(metadata) # event_caller needs session_id — it calls back to a specific # websocket session (used by direct tools, pyodide code interpreter). if metadata.get('session_id') and metadata.get('chat_id') and metadata.get('message_id'): - event_caller = get_event_call(metadata) + event_caller = await get_event_call(metadata) return event_emitter, event_caller -def build_chat_response_context(request, form_data, user, model, metadata, tasks, events): - event_emitter, event_caller = get_event_emitter_and_caller(metadata) +async def build_chat_response_context(request, form_data, user, model, metadata, tasks, events): + event_emitter, event_caller = await get_event_emitter_and_caller(metadata) return { 'request': request, 'form_data': form_data, @@ -2862,7 +2863,7 @@ async def background_tasks_handler(ctx): messages = [] if 'chat_id' in metadata and not metadata['chat_id'].startswith('local:'): - messages_map = Chats.get_messages_map_by_chat_id(metadata['chat_id']) + messages_map = await Chats.get_messages_map_by_chat_id(metadata['chat_id']) message = messages_map.get(metadata['message_id']) if messages_map else None message_list = get_message_list(messages_map, metadata['message_id']) @@ -2942,7 +2943,7 @@ async def background_tasks_handler(ctx): ) if not metadata.get('chat_id', '').startswith('local:'): - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], { @@ -2995,7 +2996,7 @@ async def background_tasks_handler(ctx): if not title: title = messages[0].get('content', user_message) - Chats.update_chat_title_by_id(metadata['chat_id'], title) + await Chats.update_chat_title_by_id(metadata['chat_id'], title) await event_emitter( { @@ -3007,7 +3008,7 @@ async def background_tasks_handler(ctx): if title == None and len(messages) == 2 and (not messages_map or len(messages_map) <= 2): title = messages[0].get('content', user_message) - Chats.update_chat_title_by_id(metadata['chat_id'], title) + await Chats.update_chat_title_by_id(metadata['chat_id'], title) await event_emitter( { @@ -3041,7 +3042,7 @@ async def background_tasks_handler(ctx): try: tags = json.loads(tags_string).get('tags', []) - Chats.update_chat_tags_by_id(metadata['chat_id'], tags, user) + await Chats.update_chat_tags_by_id(metadata['chat_id'], tags, user) await event_emitter( { @@ -3076,7 +3077,7 @@ async def non_streaming_chat_response_handler(response, ctx): else: error = str(error) - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], { @@ -3092,7 +3093,7 @@ async def non_streaming_chat_response_handler(response, ctx): ) if 'selected_model_id' in response_data: - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], { @@ -3112,7 +3113,7 @@ async def non_streaming_chat_response_handler(response, ctx): } ) - title = Chats.get_chat_title_by_id(metadata['chat_id']) + title = await Chats.get_chat_title_by_id(metadata['chat_id']) # Use output from backend if provided (OR-compliant backends), # otherwise generate from response content @@ -3143,7 +3144,7 @@ async def non_streaming_chat_response_handler(response, ctx): # Save message in the database usage = normalize_usage(response_data.get('usage', {}) or {}) - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], { @@ -3156,8 +3157,8 @@ async def non_streaming_chat_response_handler(response, ctx): ) # Send a webhook notification if the user is not active - if request.app.state.config.ENABLE_USER_WEBHOOKS and not Users.is_user_active(user.id): - webhook_url = Users.get_user_webhook_url_by_id(user.id) + if request.app.state.config.ENABLE_USER_WEBHOOKS and not await Users.is_user_active(user.id): + webhook_url = await Users.get_user_webhook_url_by_id(user.id) if webhook_url: await post_webhook( request.app.state.WEBUI_NAME, @@ -3211,8 +3212,8 @@ async def streaming_chat_response_handler(response, ctx): } filter_functions = [ - Functions.get_function_by_id(filter_id) - for filter_id in get_sorted_filter_ids(request, model, metadata.get('filter_ids', [])) + await Functions.get_function_by_id(filter_id) + for filter_id in await get_sorted_filter_ids(request, model, metadata.get('filter_ids', [])) ] # Standard streaming response handler @@ -3447,7 +3448,7 @@ async def streaming_chat_response_handler(response, ctx): return output, end_flag - message = Chats.get_message_by_id_and_message_id(metadata['chat_id'], metadata['message_id']) + message = await Chats.get_message_by_id_and_message_id(metadata['chat_id'], metadata['message_id']) tool_calls = [] @@ -3509,7 +3510,7 @@ async def streaming_chat_response_handler(response, ctx): ) # Save message in the database - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], { @@ -3579,7 +3580,7 @@ async def streaming_chat_response_handler(response, ctx): if 'selected_model_id' in data: model_id = data['selected_model_id'] - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], { @@ -3645,7 +3646,7 @@ async def streaming_chat_response_handler(response, ctx): error = data.get('error', {}) if error: try: - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], { @@ -3762,10 +3763,10 @@ async def streaming_chat_response_handler(response, ctx): } ) - image_urls = get_image_urls(delta.get('images', []), request, metadata, user) + image_urls = await get_image_urls(delta.get('images', []), request, metadata, user) if image_urls: image_file_list = [{'type': 'image', 'url': url} for url in image_urls] - message_files = Chats.add_message_files_by_id_and_message_id( + message_files = await Chats.add_message_files_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], image_file_list, @@ -3847,7 +3848,7 @@ async def streaming_chat_response_handler(response, ctx): ) if ENABLE_CHAT_RESPONSE_BASE64_IMAGE_URL_CONVERSION: - value = convert_markdown_base64_images( + value = await convert_markdown_base64_images( request, value, { @@ -3963,7 +3964,7 @@ async def streaming_chat_response_handler(response, ctx): if ENABLE_REALTIME_CHAT_SAVE: # Save message in the database - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], { @@ -4184,7 +4185,7 @@ async def streaming_chat_response_handler(response, ctx): ) else: - tool_function = get_updated_tool_function( + tool_function = await get_updated_tool_function( function=tool['callable'], extra_params={ '__messages__': form_data.get('messages', []), @@ -4197,7 +4198,7 @@ async def streaming_chat_response_handler(response, ctx): except Exception as e: tool_result = str(e) - tool_result, tool_result_files, tool_result_embeds = process_tool_result( + tool_result, tool_result_files, tool_result_embeds = await process_tool_result( request, tool_function_name, tool_result, @@ -4487,7 +4488,7 @@ async def streaming_chat_response_handler(response, ctx): BLOCKED_MODULES = {CODE_INTERPRETER_BLOCKED_MODULES} _real_import = builtins.__import__ - def restricted_import(name, globals=None, locals=None, fromlist=(), level=0): + async def restricted_import(name, globals=None, locals=None, fromlist=(), level=0): if name.split('.')[0] in BLOCKED_MODULES: importer_name = globals.get('__name__') if globals else None if importer_name == '__main__': @@ -4541,7 +4542,7 @@ async def streaming_chat_response_handler(response, ctx): stdoutLines = stdout.split('\n') for idx, line in enumerate(stdoutLines): if re.match(r'data:image/\w+;base64', line): - image_url = get_image_url_from_base64( + image_url = await get_image_url_from_base64( request, line, metadata, @@ -4558,7 +4559,7 @@ async def streaming_chat_response_handler(response, ctx): resultLines = result.split('\n') for idx, line in enumerate(resultLines): if re.match(r'data:image/\w+;base64', line): - image_url = get_image_url_from_base64( + image_url = await get_image_url_from_base64( request, line, metadata, @@ -4623,7 +4624,7 @@ async def streaming_chat_response_handler(response, ctx): if item.get('status') == 'in_progress': item['status'] = 'completed' - title = Chats.get_chat_title_by_id(metadata['chat_id']) + title = await Chats.get_chat_title_by_id(metadata['chat_id']) data = { 'done': True, 'content': serialize_output(output), @@ -4634,7 +4635,7 @@ async def streaming_chat_response_handler(response, ctx): if not ENABLE_REALTIME_CHAT_SAVE: # Save message in the database - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], { @@ -4645,21 +4646,21 @@ async def streaming_chat_response_handler(response, ctx): }, ) elif usage: - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], {'done': True, 'usage': usage}, ) else: - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], {'done': True}, ) # Send a webhook notification if the user is not active - if request.app.state.config.ENABLE_USER_WEBHOOKS and not Users.is_user_active(user.id): - webhook_url = Users.get_user_webhook_url_by_id(user.id) + if request.app.state.config.ENABLE_USER_WEBHOOKS and not await Users.is_user_active(user.id): + webhook_url = await Users.get_user_webhook_url_by_id(user.id) if webhook_url: await post_webhook( request.app.state.WEBUI_NAME, @@ -4687,7 +4688,7 @@ async def streaming_chat_response_handler(response, ctx): if not ENABLE_REALTIME_CHAT_SAVE: # Save message in the database - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], { @@ -4697,7 +4698,7 @@ async def streaming_chat_response_handler(response, ctx): }, ) else: - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], {'done': True}, diff --git a/backend/open_webui/utils/models.py b/backend/open_webui/utils/models.py index b57c74744c..c8ebc190e5 100644 --- a/backend/open_webui/utils/models.py +++ b/backend/open_webui/utils/models.py @@ -130,13 +130,13 @@ async def get_all_models(request, refresh: bool = False, user: UserModel = None) ] models = models + arena_models - global_action_ids = {function.id for function in Functions.get_global_action_functions()} - enabled_action_ids = {function.id for function in Functions.get_functions_by_type('action', active_only=True)} + global_action_ids = {function.id for function in await Functions.get_global_action_functions()} + enabled_action_ids = {function.id for function in await Functions.get_functions_by_type('action', active_only=True)} - global_filter_ids = {function.id for function in Functions.get_global_filter_functions()} - enabled_filter_ids = {function.id for function in Functions.get_functions_by_type('filter', active_only=True)} + global_filter_ids = {function.id for function in await Functions.get_global_filter_functions()} + enabled_filter_ids = {function.id for function in await Functions.get_functions_by_type('filter', active_only=True)} - custom_models = Models.get_all_models() + custom_models = await Models.get_all_models() # Single O(1) lookup: Ollama base names first, then exact IDs (exact wins). base_model_lookup = {} @@ -278,14 +278,14 @@ async def get_all_models(request, refresh: bool = False, user: UserModel = None) all_function_ids.update(global_action_ids) all_function_ids.update(global_filter_ids) - functions_by_id = {f.id: f for f in Functions.get_functions_by_ids(list(all_function_ids))} + functions_by_id = {f.id: f for f in await Functions.get_functions_by_ids(list(all_function_ids))} # Pre-warm the function module cache once per unique function ID. # This ensures each function's DB freshness check runs exactly once, # not once per (model × function) pair. for function_id in all_function_ids: try: - get_function_module_from_cache(request, function_id) + await get_function_module_from_cache(request, function_id) except Exception as e: log.info(f'Failed to load function module for {function_id}: {e}') @@ -312,7 +312,7 @@ async def get_all_models(request, refresh: bool = False, user: UserModel = None) # Batch-fetch all function valves in one query to avoid N+1 DB hits # inside get_action_priority (previously called per action × per model). - all_function_valves = Functions.get_function_valves_by_ids(list(all_function_ids)) + all_function_valves = await Functions.get_function_valves_by_ids(list(all_function_ids)) def get_action_priority(action_id): try: @@ -377,11 +377,11 @@ async def get_all_models(request, refresh: bool = False, user: UserModel = None) return models -def check_model_access(user, model, db=None): +async def check_model_access(user, model, db=None): if model.get('arena'): meta = model.get('info', {}).get('meta', {}) access_grants = meta.get('access_grants', []) - if not has_access( + if not await has_access( user.id, permission='read', access_grants=access_grants, @@ -389,12 +389,12 @@ def check_model_access(user, model, db=None): ): raise Exception('Model not found') else: - model_info = Models.get_model_by_id(model.get('id'), db=db) + model_info = await Models.get_model_by_id(model.get('id'), db=db) if not model_info: raise Exception('Model not found') elif not ( user.id == model_info.user_id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='model', resource_id=model_info.id, @@ -405,7 +405,7 @@ def check_model_access(user, model, db=None): raise Exception('Model not found') -def get_filtered_models(models, user, db=None): +async def get_filtered_models(models, user, db=None): # Filter out models that the user does not have access to if ( user.role == 'user' or (user.role == 'admin' and not BYPASS_ADMIN_ACCESS_CONTROL) @@ -418,10 +418,10 @@ def get_filtered_models(models, user, db=None): if info: model_infos[model['id']] = info - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id, db=db)} + user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)} # Batch-fetch accessible resource IDs in a single query instead of N has_access calls - accessible_model_ids = AccessGrants.get_accessible_resource_ids( + accessible_model_ids = await AccessGrants.get_accessible_resource_ids( user_id=user.id, resource_type='model', resource_ids=list(model_infos.keys()), @@ -435,7 +435,7 @@ def get_filtered_models(models, user, db=None): if model.get('arena'): meta = model.get('info', {}).get('meta', {}) access_grants = meta.get('access_grants', []) - if has_access( + if await has_access( user.id, permission='read', access_grants=access_grants, diff --git a/backend/open_webui/utils/oauth.py b/backend/open_webui/utils/oauth.py index 100df7a219..535adca5ec 100644 --- a/backend/open_webui/utils/oauth.py +++ b/backend/open_webui/utils/oauth.py @@ -700,7 +700,7 @@ class OAuthClientManager: """ try: # Get the OAuth session - session = OAuthSessions.get_session_by_provider_and_user_id(client_id, user_id) + session = await OAuthSessions.get_session_by_provider_and_user_id(client_id, user_id) if not session: log.warning(f'No OAuth session found for user {user_id}, client_id {client_id}') return None @@ -714,7 +714,7 @@ class OAuthClientManager: log.warning( f'Token refresh failed for user {user_id}, client_id {session.provider}, deleting session {session.id}' ) - OAuthSessions.delete_session_by_id(session.id) + await OAuthSessions.delete_session_by_id(session.id) return None return session.token @@ -738,7 +738,7 @@ class OAuthClientManager: if refreshed_token: # Update the session with new token data - session = OAuthSessions.update_session_by_id(session.id, refreshed_token) + session = await OAuthSessions.update_session_by_id(session.id, refreshed_token) log.info(f'Successfully refreshed token for session {session.id}') return session.token else: @@ -884,12 +884,12 @@ class OAuthClientManager: token['expires_at'] = datetime.now().timestamp() + token['expires_in'] # Clean up any existing sessions for this user/client_id first - sessions = OAuthSessions.get_sessions_by_user_id(user_id) + sessions = await OAuthSessions.get_sessions_by_user_id(user_id) for session in sessions: if session.provider == client_id: - OAuthSessions.delete_session_by_id(session.id) + await OAuthSessions.delete_session_by_id(session.id) - session = OAuthSessions.create_session( + session = await OAuthSessions.create_session( user_id=user_id, provider=client_id, token=token, @@ -963,7 +963,7 @@ class OAuthManager: """ try: # Get the OAuth session - session = OAuthSessions.get_session_by_id_and_user_id(session_id, user_id) + session = await OAuthSessions.get_session_by_id_and_user_id(session_id, user_id) if not session: log.warning(f'No OAuth session found for user {user_id}, session {session_id}') return None @@ -977,7 +977,7 @@ class OAuthManager: log.warning( f'Token refresh failed for user {user_id}, provider {session.provider}, deleting session {session.id}' ) - OAuthSessions.delete_session_by_id(session.id) + await OAuthSessions.delete_session_by_id(session.id) return None return session.token @@ -1002,7 +1002,7 @@ class OAuthManager: if refreshed_token: # Update the session with new token data - session = OAuthSessions.update_session_by_id(session.id, refreshed_token) + session = await OAuthSessions.update_session_by_id(session.id, refreshed_token) log.info(f'Successfully refreshed token for session {session.id}') return session.token else: @@ -1102,8 +1102,8 @@ class OAuthManager: log.error(f'Exception during token refresh for provider {provider}: {e}') return None - def get_user_role(self, user, user_data): - user_count = Users.get_num_users() + async def get_user_role(self, user, user_data): + user_count = await Users.get_num_users() if user and user_count == 1: # If the user is the only user, assign the role "admin" - actually repairs role for single user on login log.debug('Assigning the only user the admin role') @@ -1188,7 +1188,7 @@ class OAuthManager: return role - def update_user_groups(self, user, user_data, default_permissions, db=None): + async def update_user_groups(self, user, user_data, default_permissions, db=None): log.debug('Running OAUTH Group management') oauth_claim = auth_manager_config.OAUTH_GROUPS_CLAIM @@ -1217,8 +1217,8 @@ class OAuthManager: else: user_oauth_groups = [] - user_current_groups: list[GroupModel] = Groups.get_groups_by_member_id(user.id, db=db) - all_available_groups: list[GroupModel] = Groups.get_all_groups(db=db) + user_current_groups: list[GroupModel] = await Groups.get_groups_by_member_id(user.id, db=db) + all_available_groups: list[GroupModel] = await Groups.get_all_groups(db=db) # Create groups if they don't exist and creation is enabled if auth_manager_config.ENABLE_OAUTH_GROUP_CREATION: @@ -1226,7 +1226,7 @@ class OAuthManager: all_group_names = {g.name for g in all_available_groups} groups_created = False # Determine creator ID: Prefer admin, fallback to current user if no admin exists - admin_user = Users.get_super_admin_user() + admin_user = await Users.get_super_admin_user() creator_id = admin_user.id if admin_user else user.id log.debug(f'Using creator ID {creator_id} for potential group creation.') @@ -1241,7 +1241,7 @@ class OAuthManager: data={'config': {'share': auth_manager_config.OAUTH_GROUP_DEFAULT_SHARE}}, ) # Use determined creator ID (admin or fallback to current user) - created_group = Groups.insert_new_group(creator_id, new_group_form, db=db) + created_group = await Groups.insert_new_group(creator_id, new_group_form, db=db) if created_group: log.info( f"Successfully created group '{group_name}' with ID {created_group.id} using creator ID {creator_id}" @@ -1256,7 +1256,7 @@ class OAuthManager: # Refresh the list of all available groups if any were created if groups_created: - all_available_groups = Groups.get_all_groups(db=db) + all_available_groups = await Groups.get_all_groups(db=db) log.debug('Refreshed list of all available groups after creation.') log.debug(f'Oauth Groups claim: {oauth_claim}') @@ -1273,14 +1273,14 @@ class OAuthManager: ): # Remove group from user log.debug(f'Removing user from group {group_model.name} as it is no longer in their oauth groups') - Groups.remove_users_from_group(group_model.id, [user.id], db=db) + await Groups.remove_users_from_group(group_model.id, [user.id], db=db) # In case a group is created, but perms are never assigned to the group by hitting "save" group_permissions = group_model.permissions if not group_permissions: group_permissions = default_permissions - Groups.update_group_by_id( + await Groups.update_group_by_id( id=group_model.id, form_data=GroupUpdateForm( name=group_model.name, @@ -1302,14 +1302,14 @@ class OAuthManager: # Add user to group log.debug(f'Adding user to group {group_model.name} as it was found in their oauth groups') - Groups.add_users_to_group(group_model.id, [user.id], db=db) + await Groups.add_users_to_group(group_model.id, [user.id], db=db) # In case a group is created, but perms are never assigned to the group by hitting "save" group_permissions = group_model.permissions if not group_permissions: group_permissions = default_permissions - Groups.update_group_by_id( + await Groups.update_group_by_id( id=group_model.id, form_data=GroupUpdateForm( name=group_model.name, @@ -1490,20 +1490,20 @@ class OAuthManager: raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED) # Check if the user exists - user = Users.get_user_by_oauth_sub(provider, sub, db=db) + user = await Users.get_user_by_oauth_sub(provider, sub, db=db) if not user: # If the user does not exist, check if merging is enabled if auth_manager_config.OAUTH_MERGE_ACCOUNTS_BY_EMAIL: # Check if the user exists by email - user = Users.get_user_by_email(email, db=db) + user = await Users.get_user_by_email(email, db=db) if user: # Update the user with the new oauth sub - Users.update_user_oauth_by_id(user.id, provider, sub, db=db) + await Users.update_user_oauth_by_id(user.id, provider, sub, db=db) if user: - determined_role = self.get_user_role(user, user_data) + determined_role = await self.get_user_role(user, user_data) if user.role != determined_role: - Users.update_user_role_by_id(user.id, determined_role, db=db) + await Users.update_user_role_by_id(user.id, determined_role, db=db) # Update the user object in memory as well, # to avoid problems with the ENABLE_OAUTH_GROUP_MANAGEMENT check below user.role = determined_role @@ -1513,7 +1513,7 @@ class OAuthManager: if username_claim: new_name = user_data.get(username_claim) if new_name and new_name != user.name: - Users.update_user_by_id(user.id, {'name': new_name}, db=db) + await Users.update_user_by_id(user.id, {'name': new_name}, db=db) user.name = new_name log.debug(f'Updated name for user {user.email}') @@ -1522,13 +1522,13 @@ class OAuthManager: if email_claim: new_email = user_data.get(email_claim) if new_email and new_email.lower() != user.email.lower(): - existing_user = Users.get_user_by_email(new_email, db=db) + existing_user = await Users.get_user_by_email(new_email, db=db) if existing_user: log.error( f'Cannot update email to {new_email} for user {user.id} because it is already taken.' ) else: - Auths.update_email_by_id(user.id, new_email.lower(), db=db) + await Auths.update_email_by_id(user.id, new_email.lower(), db=db) user.email = new_email.lower() log.debug(f'Updated email for user {user.id}') @@ -1544,13 +1544,13 @@ class OAuthManager: new_picture_url, token.get('access_token') ) if processed_picture_url != user.profile_image_url: - Users.update_user_profile_image_url_by_id(user.id, processed_picture_url, db=db) + await Users.update_user_profile_image_url_by_id(user.id, processed_picture_url, db=db) log.debug(f'Updated profile picture for user {user.email}') else: # If the user does not exist, check if signups are enabled if auth_manager_config.ENABLE_OAUTH_SIGNUP: # Check if an existing user with the same email already exists - existing_user = Users.get_user_by_email(email, db=db) + existing_user = await Users.get_user_by_email(email, db=db) if existing_user: raise HTTPException(400, detail=ERROR_MESSAGES.EMAIL_TAKEN) @@ -1570,12 +1570,12 @@ class OAuthManager: log.warning('Username claim is missing, using email as name') name = email - user = Auths.insert_new_auth( + user = await Auths.insert_new_auth( email=email, password=get_password_hash(str(uuid.uuid4())), # Random password, not used name=name, profile_image_url=picture_url, - role=self.get_user_role(None, user_data), + role=await self.get_user_role(None, user_data), oauth=oauth_data, db=db, ) @@ -1586,9 +1586,9 @@ class OAuthManager: # Atomically check if this is the only user *after* the # insert to avoid TOCTOU race on first-user registration. # Matches signup_handler pattern. - if Users.get_num_users(db=db) == 1: - Users.update_user_role_by_id(user.id, 'admin', db=db) - user = Users.get_user_by_id(user.id, db=db) + if await Users.get_num_users(db=db) == 1: + await Users.update_user_role_by_id(user.id, 'admin', db=db) + user = await Users.get_user_by_id(user.id, db=db) if auth_manager_config.WEBHOOK_URL: await post_webhook( @@ -1602,7 +1602,7 @@ class OAuthManager: }, ) - apply_default_group_assignment(request.app.state.config.DEFAULT_GROUP_ID, user.id, db=db) + await apply_default_group_assignment(request.app.state.config.DEFAULT_GROUP_ID, user.id, db=db) else: raise HTTPException( @@ -1615,7 +1615,7 @@ class OAuthManager: expires_delta=parse_duration(auth_manager_config.JWT_EXPIRES_IN), ) if auth_manager_config.ENABLE_OAUTH_GROUP_MANAGEMENT: - self.update_user_groups( + await self.update_user_groups( user=user, user_data=user_data, default_permissions=request.app.state.config.USER_PERMISSIONS, @@ -1675,7 +1675,7 @@ class OAuthManager: # Enforce max concurrent sessions per user/provider to prevent # unbounded growth while allowing multi-device usage - sessions = OAuthSessions.get_sessions_by_user_id(user.id, db=db) + sessions = await OAuthSessions.get_sessions_by_user_id(user.id, db=db) provider_sessions = sorted( [session for session in sessions if session.provider == provider], key=lambda session: session.created_at, @@ -1684,9 +1684,9 @@ class OAuthManager: # Keep the newest sessions up to the limit, prune the rest if len(provider_sessions) >= OAUTH_MAX_SESSIONS_PER_USER: for old_session in provider_sessions[OAUTH_MAX_SESSIONS_PER_USER - 1 :]: - OAuthSessions.delete_session_by_id(old_session.id, db=db) + await OAuthSessions.delete_session_by_id(old_session.id, db=db) - session = OAuthSessions.create_session( + session = await OAuthSessions.create_session( user_id=user.id, provider=provider, token=token, @@ -1847,7 +1847,7 @@ class OAuthManager: # 8. Identify users to log out users_to_logout = [] if sub: - user = Users.get_user_by_oauth_sub(matched_provider, sub, db=db) + user = await Users.get_user_by_oauth_sub(matched_provider, sub, db=db) if user: users_to_logout.append(user) @@ -1868,9 +1868,9 @@ class OAuthManager: revoked_count = 0 for user in users_to_logout: - sessions = OAuthSessions.get_sessions_by_user_id(user.id, db=db) + sessions = await OAuthSessions.get_sessions_by_user_id(user.id, db=db) for oauth_session in sessions: - OAuthSessions.delete_session_by_id(oauth_session.id, db=db) + await OAuthSessions.delete_session_by_id(oauth_session.id, db=db) if redis: revocation_key = f'{REDIS_KEY_PREFIX}:auth:user:{user.id}:revoked_at' diff --git a/backend/open_webui/utils/plugin.py b/backend/open_webui/utils/plugin.py index 46622e21ae..84671bbd3b 100644 --- a/backend/open_webui/utils/plugin.py +++ b/backend/open_webui/utils/plugin.py @@ -199,16 +199,16 @@ def replace_imports(content): # May the intent of the one who wrote it survive every # import and transformation, as a deed survives the generations. -def load_tool_module_by_id(tool_id, content=None): +async def load_tool_module_by_id(tool_id, content=None): if content is None: - tool = Tools.get_tool_by_id(tool_id) + tool = await Tools.get_tool_by_id(tool_id) if not tool: raise Exception(f'Toolkit not found: {tool_id}') content = tool.content content = replace_imports(content) - Tools.update_tool_by_id(tool_id, {'content': content}) + await Tools.update_tool_by_id(tool_id, {'content': content}) else: frontmatter = extract_frontmatter(content) # Install required packages found within the frontmatter @@ -245,15 +245,15 @@ def load_tool_module_by_id(tool_id, content=None): os.unlink(temp_file.name) -def load_function_module_by_id(function_id: str, content: str | None = None): +async def load_function_module_by_id(function_id: str, content: str | None = None): if content is None: - function = Functions.get_function_by_id(function_id) + function = await Functions.get_function_by_id(function_id) if not function: raise Exception(f'Function not found: {function_id}') content = function.content content = replace_imports(content) - Functions.update_function_by_id(function_id, {'content': content}) + await Functions.update_function_by_id(function_id, {'content': content}) else: frontmatter = extract_frontmatter(content) install_frontmatter_requirements(frontmatter.get('requirements', '')) @@ -290,16 +290,16 @@ def load_function_module_by_id(function_id: str, content: str | None = None): # Cleanup by removing the module in case of error del sys.modules[module_name] - Functions.update_function_by_id(function_id, {'is_active': False}) + await Functions.update_function_by_id(function_id, {'is_active': False}) raise e finally: os.unlink(temp_file.name) -def get_tool_module_from_cache(request, tool_id, load_from_db=True): +async def get_tool_module_from_cache(request, tool_id, load_from_db=True): if load_from_db: # Always load from the database by default - tool = Tools.get_tool_by_id(tool_id) + tool = await Tools.get_tool_by_id(tool_id) if not tool: raise Exception(f'Tool not found: {tool_id}') content = tool.content @@ -308,7 +308,7 @@ def get_tool_module_from_cache(request, tool_id, load_from_db=True): if new_content != content: content = new_content # Update the tool content in the database - Tools.update_tool_by_id(tool_id, {'content': content}) + await Tools.update_tool_by_id(tool_id, {'content': content}) if (hasattr(request.app.state, 'TOOL_CONTENTS') and tool_id in request.app.state.TOOL_CONTENTS) and ( hasattr(request.app.state, 'TOOLS') and tool_id in request.app.state.TOOLS @@ -316,12 +316,12 @@ def get_tool_module_from_cache(request, tool_id, load_from_db=True): if request.app.state.TOOL_CONTENTS[tool_id] == content: return request.app.state.TOOLS[tool_id], None - tool_module, frontmatter = load_tool_module_by_id(tool_id, content) + tool_module, frontmatter = await load_tool_module_by_id(tool_id, content) else: if hasattr(request.app.state, 'TOOLS') and tool_id in request.app.state.TOOLS: return request.app.state.TOOLS[tool_id], None - tool_module, frontmatter = load_tool_module_by_id(tool_id) + tool_module, frontmatter = await load_tool_module_by_id(tool_id) if not hasattr(request.app.state, 'TOOLS'): request.app.state.TOOLS = {} @@ -335,13 +335,13 @@ def get_tool_module_from_cache(request, tool_id, load_from_db=True): return tool_module, frontmatter -def get_function_module_from_cache(request, function_id, load_from_db=True): +async def get_function_module_from_cache(request, function_id, load_from_db=True): if load_from_db: # Always load from the database by default # This is useful for hooks like "inlet" or "outlet" where the content might change # and we want to ensure the latest content is used. - function = Functions.get_function_by_id(function_id) + function = await Functions.get_function_by_id(function_id) if not function: raise Exception(f'Function not found: {function_id}') content = function.content @@ -350,7 +350,7 @@ def get_function_module_from_cache(request, function_id, load_from_db=True): if new_content != content: content = new_content # Update the function content in the database - Functions.update_function_by_id(function_id, {'content': content}) + await Functions.update_function_by_id(function_id, {'content': content}) if ( hasattr(request.app.state, 'FUNCTION_CONTENTS') and function_id in request.app.state.FUNCTION_CONTENTS @@ -358,7 +358,7 @@ def get_function_module_from_cache(request, function_id, load_from_db=True): if request.app.state.FUNCTION_CONTENTS[function_id] == content: return request.app.state.FUNCTIONS[function_id], None, None - function_module, function_type, frontmatter = load_function_module_by_id(function_id, content) + function_module, function_type, frontmatter = await load_function_module_by_id(function_id, content) else: # Load from cache (e.g. "stream" hook) # This is useful for performance reasons @@ -366,7 +366,7 @@ def get_function_module_from_cache(request, function_id, load_from_db=True): if hasattr(request.app.state, 'FUNCTIONS') and function_id in request.app.state.FUNCTIONS: return request.app.state.FUNCTIONS[function_id], None, None - function_module, function_type, frontmatter = load_function_module_by_id(function_id) + function_module, function_type, frontmatter = await load_function_module_by_id(function_id) if not hasattr(request.app.state, 'FUNCTIONS'): request.app.state.FUNCTIONS = {} @@ -404,7 +404,7 @@ def install_frontmatter_requirements(requirements: str): log.info('No requirements found in frontmatter.') -def install_tool_and_function_dependencies(): +async def install_tool_and_function_dependencies(): """ Install all dependencies for all admin tools and active functions. @@ -412,8 +412,8 @@ 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. """ - function_list = Functions.get_functions(active_only=True) - tool_list = Tools.get_tools() + function_list = await Functions.get_functions(active_only=True) + tool_list = await Tools.get_tools() all_dependencies = '' try: diff --git a/backend/open_webui/utils/redis.py b/backend/open_webui/utils/redis.py index 61a4d74c7c..7a114393b0 100644 --- a/backend/open_webui/utils/redis.py +++ b/backend/open_webui/utils/redis.py @@ -38,7 +38,7 @@ class SentinelRedisProxy: def _master(self): return self._sentinel.master_for(self._service, **self._kw) - def __getattr__(self, item): + async def __getattr__(self, item): master = self._master() orig_attr = getattr(master, item) diff --git a/backend/open_webui/utils/telemetry/metrics.py b/backend/open_webui/utils/telemetry/metrics.py index 4c43de3342..26216b6ca4 100644 --- a/backend/open_webui/utils/telemetry/metrics.py +++ b/backend/open_webui/utils/telemetry/metrics.py @@ -124,16 +124,16 @@ def setup_metrics(app: FastAPI, resource: Resource) -> None: unit='ms', ) - def observe_active_users( + async def observe_active_users( options: metrics.CallbackOptions, ) -> Sequence[metrics.Observation]: return [ metrics.Observation( - value=Users.get_active_user_count(), + value=await Users.get_active_user_count(), ) ] - def observe_total_registered_users( + async def observe_total_registered_users( options: metrics.CallbackOptions, ) -> Sequence[metrics.Observation]: # IMPORTANT: Use get_num_users() for efficient COUNT(*) query. @@ -141,7 +141,7 @@ def setup_metrics(app: FastAPI, resource: Resource) -> None: # causing connection pool exhaustion on high-latency databases (e.g., Aurora). return [ metrics.Observation( - value=Users.get_num_users() or 0, + value=await Users.get_num_users() or 0, ) ] @@ -159,10 +159,10 @@ def setup_metrics(app: FastAPI, resource: Resource) -> None: callbacks=[observe_active_users], ) - def observe_users_active_today( + async def observe_users_active_today( options: metrics.CallbackOptions, ) -> Sequence[metrics.Observation]: - return [metrics.Observation(value=Users.get_num_users_active_today())] + return [metrics.Observation(value=await Users.get_num_users_active_today())] meter.create_observable_gauge( name='webui.users.active.today', diff --git a/backend/open_webui/utils/tools.py b/backend/open_webui/utils/tools.py index 2223f202a0..a44fe69ab8 100644 --- a/backend/open_webui/utils/tools.py +++ b/backend/open_webui/utils/tools.py @@ -101,7 +101,7 @@ log = logging.getLogger(__name__) # Let no function be called without need, and let what # it yields justify the cost of running it. -def get_async_tool_function_and_apply_extra_params(function: Callable, extra_params: dict) -> Callable[..., Awaitable]: +async def get_async_tool_function_and_apply_extra_params(function: Callable, extra_params: dict) -> Callable[..., Awaitable]: sig = inspect.signature(function) extra_params = {k: v for k, v in extra_params.items() if k in sig.parameters} partial_func = partial(function, **extra_params) @@ -138,13 +138,13 @@ def get_async_tool_function_and_apply_extra_params(function: Callable, extra_par return new_function -def get_updated_tool_function(function: Callable, extra_params: dict): +async def get_updated_tool_function(function: Callable, extra_params: dict): # Get the original function and merge updated params __function__ = getattr(function, '__function__', None) __extra_params__ = getattr(function, '__extra_params__', None) if __function__ is not None and __extra_params__ is not None: - return get_async_tool_function_and_apply_extra_params( + return await get_async_tool_function_and_apply_extra_params( __function__, {**__extra_params__, **extra_params}, ) @@ -160,16 +160,16 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr tools_dict = {} # Get user's group memberships for access control checks - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)} + user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)} for tool_id in tool_ids: - tool = Tools.get_tool_by_id(tool_id) + tool = await Tools.get_tool_by_id(tool_id) if tool: # Check access control for local tools if ( not (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL) and tool.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='tool', resource_id=tool.id, @@ -182,7 +182,7 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr module = request.app.state.TOOLS.get(tool_id, None) if module is None: - module, _ = load_tool_module_by_id(tool_id) + module, _ = await load_tool_module_by_id(tool_id) request.app.state.TOOLS[tool_id] = module __user__ = { @@ -191,11 +191,11 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr # Set valves for the tool if hasattr(module, 'valves') and hasattr(module, 'Valves'): - valves = Tools.get_tool_valves_by_id(tool_id) or {} + valves = await Tools.get_tool_valves_by_id(tool_id) or {} module.valves = module.Valves(**valves) if hasattr(module, 'UserValves'): __user__['valves'] = module.UserValves( # type: ignore - **Tools.get_user_valves_by_id_and_user_id(tool_id, user.id) + **await Tools.get_user_valves_by_id_and_user_id(tool_id, user.id) ) for spec in tool.specs: @@ -213,7 +213,7 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr # convert to function that takes only model params and inserts custom params function_name = spec['name'] tool_function = getattr(module, function_name) - callable = get_async_tool_function_and_apply_extra_params( + callable = await get_async_tool_function_and_apply_extra_params( tool_function, { **extra_params, @@ -285,7 +285,7 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr tool_server_connection = connections[tool_server_idx] # Check access control for tool server - if not has_connection_access(user, tool_server_connection, user_group_ids): + if not await has_connection_access(user, tool_server_connection, user_group_ids): log.warning(f'Access denied to tool server {server_id} for user {user.id}') continue @@ -339,7 +339,7 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr if metadata and metadata.get('message_id'): headers[FORWARD_SESSION_INFO_HEADER_MESSAGE_ID] = metadata.get('message_id') - def make_tool_function(function_name, tool_server_data, headers): + async def make_tool_function(function_name, tool_server_data, headers): async def tool_function(**kwargs): return await execute_tool_server( url=tool_server_data['url'], @@ -352,9 +352,9 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr return tool_function - tool_function = make_tool_function(function_name, tool_server_data, headers) + tool_function = await make_tool_function(function_name, tool_server_data, headers) - callable = get_async_tool_function_and_apply_extra_params( + callable = await get_async_tool_function_and_apply_extra_params( tool_function, {}, ) @@ -381,7 +381,7 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr return tools_dict -def get_builtin_tools( +async def get_builtin_tools( request: Request, extra_params: dict, features: dict = None, model: dict = None ) -> dict[str, dict]: """ @@ -406,10 +406,10 @@ def get_builtin_tools( # Helper to check user-level feature permission (admins always pass) user = extra_params.get('__user__', {}) - def has_user_permission(feature_key: str) -> bool: + async def has_user_permission(feature_key: str) -> bool: if user.get('role') == 'admin': return True - return has_permission( + return await has_permission( user.get('id', ''), f'features.{feature_key}', request.app.state.config.USER_PERMISSIONS, @@ -461,7 +461,7 @@ def get_builtin_tools( if ( is_builtin_tool_enabled('memory') and (features.get('memory') or get_model_capability('memory', False)) - and has_user_permission('memories') + and await has_user_permission('memories') ): builtin_functions.extend( [ @@ -479,7 +479,7 @@ def get_builtin_tools( and getattr(request.app.state.config, 'ENABLE_WEB_SEARCH', False) and get_model_capability('web_search') and features.get('web_search') - and has_user_permission('web_search') + and await has_user_permission('web_search') ): builtin_functions.extend([search_web, fetch_url]) @@ -489,7 +489,7 @@ def get_builtin_tools( and getattr(request.app.state.config, 'ENABLE_IMAGE_GENERATION', False) and get_model_capability('image_generation') and features.get('image_generation') - and has_user_permission('image_generation') + and await has_user_permission('image_generation') ): builtin_functions.append(generate_image) if ( @@ -497,7 +497,7 @@ def get_builtin_tools( and getattr(request.app.state.config, 'ENABLE_IMAGE_EDIT', False) and get_model_capability('image_generation') and features.get('image_generation') - and has_user_permission('image_generation') + and await has_user_permission('image_generation') ): builtin_functions.append(edit_image) @@ -507,7 +507,7 @@ def get_builtin_tools( and getattr(request.app.state.config, 'ENABLE_CODE_INTERPRETER', True) and get_model_capability('code_interpreter') and features.get('code_interpreter') - and has_user_permission('code_interpreter') + and await has_user_permission('code_interpreter') ): builtin_functions.append(execute_code) @@ -515,7 +515,7 @@ def get_builtin_tools( if ( is_builtin_tool_enabled('notes') and getattr(request.app.state.config, 'ENABLE_NOTES', False) - and has_user_permission('notes') + and await has_user_permission('notes') ): builtin_functions.extend([search_notes, view_note, write_note, replace_note_content]) @@ -523,7 +523,7 @@ def get_builtin_tools( if ( is_builtin_tool_enabled('channels') and getattr(request.app.state.config, 'ENABLE_CHANNELS', False) - and has_user_permission('channels') + and await has_user_permission('channels') ): builtin_functions.extend( [ @@ -543,11 +543,11 @@ def get_builtin_tools( builtin_functions.append(tasks) # Automation tools - create and manage scheduled automations from chat - if is_builtin_tool_enabled('automations') and has_user_permission('automations'): + if is_builtin_tool_enabled('automations') and await has_user_permission('automations'): builtin_functions.extend([create_automation, update_automation, list_automations, toggle_automation, delete_automation]) for func in builtin_functions: - callable = get_async_tool_function_and_apply_extra_params( + callable = await get_async_tool_function_and_apply_extra_params( func, { '__request__': request, @@ -1024,8 +1024,8 @@ async def get_terminal_tools( log.warning(f'Terminal server not found: {terminal_id}') return {} - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)} - if not has_connection_access(user, connection, user_group_ids): + user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)} + if not await has_connection_access(user, connection, user_group_ids): log.warning(f'Access denied to terminal {terminal_id} for user {user.id}') return {} @@ -1077,7 +1077,7 @@ async def get_terminal_tools( tool_spec.get('description', '') + f'\n\nThe current working directory is: {terminal_cwd}' ) - def make_tool_function(fn_name, srv_data, hdrs, cks): + async def make_tool_function(fn_name, srv_data, hdrs, cks): async def tool_function(**kwargs): return await execute_tool_server( url=srv_data['url'], @@ -1090,8 +1090,8 @@ async def get_terminal_tools( return tool_function - tool_function = make_tool_function(function_name, server_data, headers, cookies) - callable = get_async_tool_function_and_apply_extra_params(tool_function, {}) + tool_function = await make_tool_function(function_name, server_data, headers, cookies) + callable = await get_async_tool_function_and_apply_extra_params(tool_function, {}) tools_dict[function_name] = { 'tool_id': f'terminal:{terminal_id}', diff --git a/backend/requirements-min.txt b/backend/requirements-min.txt index 13bd199f08..b48006db20 100644 --- a/backend/requirements-min.txt +++ b/backend/requirements-min.txt @@ -26,6 +26,8 @@ httpx[socks,http2,zstd,cli,brotli]==0.28.1 starsessions[redis]==2.2.1 sqlalchemy==2.0.48 +aiosqlite==0.21.0 +asyncpg==0.30.0 alembic==1.18.4 peewee==3.19.0 peewee-migrate==1.14.3 diff --git a/backend/requirements.txt b/backend/requirements.txt index a9275beaf3..25265d0631 100644 --- a/backend/requirements.txt +++ b/backend/requirements.txt @@ -24,6 +24,8 @@ starsessions[redis]==2.2.1 python-mimeparse==2.0.0 sqlalchemy==2.0.48 +aiosqlite==0.21.0 +asyncpg==0.30.0 alembic==1.18.4 peewee==3.19.0 peewee-migrate==1.14.3 From d40f31982be3eed37e55e3f67b1eea9a5dc8c525 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Sun, 12 Apr 2026 14:24:08 -0500 Subject: [PATCH 32/67] refac --- backend/open_webui/routers/tools.py | 37 ++++++++++++++--------------- 1 file changed, 18 insertions(+), 19 deletions(-) diff --git a/backend/open_webui/routers/tools.py b/backend/open_webui/routers/tools.py index c61ef7752f..edf9c8b5ef 100644 --- a/backend/open_webui/routers/tools.py +++ b/backend/open_webui/routers/tools.py @@ -160,30 +160,29 @@ async def get_tools( return tools else: user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)} - tools = [ - tool - for tool in tools - if tool.user_id == user.id - or ( - has_access( + filtered_tools = [] + for tool in tools: + if tool.user_id == user.id: + filtered_tools.append(tool) + elif str(tool.id).startswith('server:'): + if await has_access( user.id, 'read', server_access_grants.get(str(tool.id), []), user_group_ids, db=db, - ) - if str(tool.id).startswith('server:') - else await AccessGrants.has_access( - user_id=user.id, - resource_type='tool', - resource_id=tool.id, - permission='read', - user_group_ids=user_group_ids, - db=db, - ) - ) - ] - return tools + ): + filtered_tools.append(tool) + elif await AccessGrants.has_access( + user_id=user.id, + resource_type='tool', + resource_id=tool.id, + permission='read', + user_group_ids=user_group_ids, + db=db, + ): + filtered_tools.append(tool) + return filtered_tools ############################ From f6b85700eafda438eb534a74585de6cf3a9194bb Mon Sep 17 00:00:00 2001 From: Classic298 <27028174+Classic298@users.noreply.github.com> Date: Sun, 12 Apr 2026 21:26:12 +0200 Subject: [PATCH 33/67] fix: gate OpenAI catch-all proxy behind ENABLE_OPENAI_API_PASSTHROUGH toggle (#23640) The catch-all /{path:path} proxy forwards any request to the upstream OpenAI-compatible API with the admin's API key and no access control. This is an intentional proxy but should be opt-in. Adds ENABLE_OPENAI_API_PASSTHROUGH env var (defaults to False). When disabled, the catch-all returns 403. No other routers (Ollama, responses) have catch-all proxies. --- backend/open_webui/env.py | 5 +++++ backend/open_webui/routers/openai.py | 12 ++++++++++-- 2 files changed, 15 insertions(+), 2 deletions(-) diff --git a/backend/open_webui/env.py b/backend/open_webui/env.py index 1695cec20a..9a16fd4ba8 100644 --- a/backend/open_webui/env.py +++ b/backend/open_webui/env.py @@ -519,6 +519,11 @@ PASSWORD_VALIDATION_HINT = os.environ.get('PASSWORD_VALIDATION_HINT', '') BYPASS_MODEL_ACCESS_CONTROL = os.environ.get('BYPASS_MODEL_ACCESS_CONTROL', 'False').lower() == 'true' +# When disabled (default), the OpenAI catch-all proxy endpoint (/{path:path}) +# is blocked. Enable only if you need direct passthrough to upstream OpenAI- +# compatible APIs for endpoints not natively handled by Open WebUI. +ENABLE_OPENAI_API_PASSTHROUGH = os.environ.get('ENABLE_OPENAI_API_PASSTHROUGH', 'False').lower() == 'true' + WEBUI_AUTH_SIGNOUT_REDIRECT_URL = os.environ.get('WEBUI_AUTH_SIGNOUT_REDIRECT_URL', None) #################################### diff --git a/backend/open_webui/routers/openai.py b/backend/open_webui/routers/openai.py index 047e56fcf6..51d4267c38 100644 --- a/backend/open_webui/routers/openai.py +++ b/backend/open_webui/routers/openai.py @@ -12,7 +12,7 @@ import requests from azure.identity import DefaultAzureCredential, get_bearer_token_provider -from fastapi import Depends, HTTPException, Request, APIRouter +from fastapi import Depends, HTTPException, Request, APIRouter, status from fastapi.responses import ( FileResponse, StreamingResponse, @@ -40,6 +40,7 @@ from open_webui.env import ( ENABLE_FORWARD_USER_INFO_HEADERS, FORWARD_SESSION_INFO_HEADER_CHAT_ID, BYPASS_MODEL_ACCESS_CONTROL, + ENABLE_OPENAI_API_PASSTHROUGH, ) from open_webui.models.users import UserModel @@ -1438,9 +1439,16 @@ async def responses( @router.api_route('/{path:path}', methods=['GET', 'POST', 'PUT', 'DELETE']) async def proxy(path: str, request: Request, user=Depends(get_verified_user)): """ - Deprecated: proxy all requests to OpenAI API + Deprecated: proxy all requests to OpenAI API. + Disabled by default. Set ENABLE_OPENAI_API_PASSTHROUGH=True to enable. """ + if not ENABLE_OPENAI_API_PASSTHROUGH: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail='Direct API passthrough is disabled. Set ENABLE_OPENAI_API_PASSTHROUGH=True to enable.', + ) + body = await request.body() # Parse JSON body to resolve model-based routing From de27a121511a31606f250ba4033490797216a0eb Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Sun, 12 Apr 2026 14:39:23 -0500 Subject: [PATCH 34/67] refac --- backend/open_webui/routers/files.py | 18 +++++++++--------- backend/open_webui/routers/images.py | 2 +- backend/open_webui/routers/knowledge.py | 13 ++++++------- backend/open_webui/tools/builtin.py | 8 ++++---- 4 files changed, 20 insertions(+), 21 deletions(-) diff --git a/backend/open_webui/routers/files.py b/backend/open_webui/routers/files.py index 8f1ee13f7f..70e0f468f9 100644 --- a/backend/open_webui/routers/files.py +++ b/backend/open_webui/routers/files.py @@ -22,7 +22,7 @@ from fastapi import ( from fastapi.responses import FileResponse, StreamingResponse from sqlalchemy.ext.asyncio import AsyncSession -from open_webui.internal.db import get_async_session, SessionLocal +from open_webui.internal.db import get_async_session, get_async_db_context from open_webui.constants import ERROR_MESSAGES from open_webui.retrieval.vector.factory import VECTOR_DB_CLIENT @@ -113,7 +113,7 @@ async def process_uploaded_file( file_path_processed = Storage.get_file(file_path) result = transcribe(request, file_path_processed, file_metadata, user) - process_file( + await process_file( request, ProcessFileForm(file_id=file_item.id, content=result.get('text', '')), user=user, @@ -122,7 +122,7 @@ async def process_uploaded_file( elif (not content_type.startswith(('image/', 'video/'))) or ( request.app.state.config.CONTENT_EXTRACTION_ENGINE == 'external' ): - process_file( + await process_file( request, ProcessFileForm(file_id=file_item.id), user=user, @@ -132,7 +132,7 @@ async def process_uploaded_file( raise Exception(f'File type {content_type} is not supported for processing') else: log.info(f'File type {file.content_type} is not provided, but trying to process anyway') - process_file( + await process_file( request, ProcessFileForm(file_id=file_item.id), user=user, @@ -151,10 +151,10 @@ async def process_uploaded_file( ) if db: - _process_handler(db) + await _process_handler(db) else: - with SessionLocal() as db_session: - _process_handler(db_session) + async with get_async_db_context() as db_session: + await _process_handler(db_session) @router.post('/', response_model=FileModelResponse) @@ -540,7 +540,7 @@ async def update_file_data_content_by_id( if file.user_id == user.id or user.role == 'admin' or await has_access_to_file(id, 'write', user, db=db): try: - process_file( + await process_file( request, ProcessFileForm(file_id=id, content=form_data.content), user=user, @@ -560,7 +560,7 @@ async def update_file_data_content_by_id( # Remove old embeddings for this file from the KB collection VECTOR_DB_CLIENT.delete(collection_name=knowledge.id, filter={'file_id': id}) # Re-add from the now-updated file-{file_id} collection - process_file( + await process_file( request, ProcessFileForm(file_id=id, collection_name=knowledge.id), user=user, diff --git a/backend/open_webui/routers/images.py b/backend/open_webui/routers/images.py index fb3bc1cec5..0d534db7f6 100644 --- a/backend/open_webui/routers/images.py +++ b/backend/open_webui/routers/images.py @@ -832,7 +832,7 @@ async def image_edits( except Exception as e: raise HTTPException(status_code=400, detail=ERROR_MESSAGES.DEFAULT(e)) - async def get_image_file_item(base64_string, param_name='image'): + def get_image_file_item(base64_string, param_name='image'): data = base64_string header, encoded = data.split(',', 1) mime_type = header.split(';')[0].lstrip('data:') diff --git a/backend/open_webui/routers/knowledge.py b/backend/open_webui/routers/knowledge.py index f6c3416c8d..3022763b49 100644 --- a/backend/open_webui/routers/knowledge.py +++ b/backend/open_webui/routers/knowledge.py @@ -2,7 +2,7 @@ from typing import List, Optional from pydantic import BaseModel from fastapi import APIRouter, Depends, HTTPException, status, Request, Query from fastapi.responses import StreamingResponse -from fastapi.concurrency import run_in_threadpool + import logging import io import zipfile @@ -319,8 +319,7 @@ async def reindex_knowledge_files( failed_files = [] for file in files: try: - await run_in_threadpool( - process_file, + await process_file( request, ProcessFileForm(file_id=file.id, collection_name=knowledge_base.id), user=user, @@ -543,7 +542,7 @@ async def update_knowledge_access_by_id( await AccessGrants.set_access_grants('knowledge', id, form_data.access_grants, db=db) return KnowledgeFilesResponse( - **await Knowledges.get_knowledge_by_id(id=id, db=db).model_dump(), + **(await Knowledges.get_knowledge_by_id(id=id, db=db)).model_dump(), files=await Knowledges.get_file_metadatas_by_id(id, db=db), ) @@ -659,7 +658,7 @@ async def add_file_to_knowledge_by_id( # Add content to the vector database try: - process_file( + await process_file( request, ProcessFileForm(file_id=form_data.file_id, collection_name=id), user=user, @@ -737,7 +736,7 @@ async def update_file_from_knowledge_by_id( # Add content to the vector database try: - process_file( + await process_file( request, ProcessFileForm(file_id=form_data.file_id, collection_name=id), user=user, @@ -962,7 +961,7 @@ async def reset_knowledge_by_id(id: str, user=Depends(get_verified_user), db: As log.debug(e) pass - knowledge = Knowledges.reset_knowledge_by_id(id=id, db=db) + knowledge = await Knowledges.reset_knowledge_by_id(id=id, db=db) return knowledge diff --git a/backend/open_webui/tools/builtin.py b/backend/open_webui/tools/builtin.py index 58af934372..25d1f84413 100644 --- a/backend/open_webui/tools/builtin.py +++ b/backend/open_webui/tools/builtin.py @@ -1213,7 +1213,7 @@ async def search_channel_messages( end_ts = end_timestamp * 1_000_000_000 if end_timestamp else None # Search messages using the model method - matching_messages = Messages.search_messages_by_channel_ids( + matching_messages = await Messages.search_messages_by_channel_ids( channel_ids=channel_ids, query=query, start_timestamp=start_ts, @@ -1274,7 +1274,7 @@ async def view_channel_message( try: user_id = __user__.get('id') - message = Messages.get_message_by_id(message_id) + message = await Messages.get_message_by_id(message_id) if not message: return json.dumps({'error': 'Message not found'}) @@ -1336,7 +1336,7 @@ async def view_channel_thread( user_id = __user__.get('id') # Get the parent message - parent_message = Messages.get_message_by_id(parent_message_id) + parent_message = await Messages.get_message_by_id(parent_message_id) if not parent_message: return json.dumps({'error': 'Message not found'}) @@ -1353,7 +1353,7 @@ async def view_channel_thread( return json.dumps({'error': 'Access denied'}) # Get all thread replies - thread_replies = Messages.get_thread_replies_by_message_id(parent_message_id) + thread_replies = await Messages.get_thread_replies_by_message_id(parent_message_id) # Build the response messages = [] From 977d638afe1b3983360abea5185c9c5a96eff45d Mon Sep 17 00:00:00 2001 From: Classic298 <27028174+Classic298@users.noreply.github.com> Date: Sun, 12 Apr 2026 23:19:38 +0200 Subject: [PATCH 35/67] fix: invalidate stale Socket.IO sessions on role change and user deletion (#23642) SESSION_POOL caches user.role at connection time and never refreshes it. When an admin demotes or deletes a user, their socket sessions retain the old cached role until voluntary disconnect, allowing continued use of admin-gated socket features (ydoc editing, channel access). Adds disconnect_user_sessions() helper that disconnects all sockets for a user ID. Called from update_user_by_id (on role change) and delete_user_by_id. The client auto-reconnects and re-authenticates with fresh DB data. --- backend/open_webui/routers/users.py | 6 ++++++ backend/open_webui/socket/main.py | 18 ++++++++++++++++++ 2 files changed, 24 insertions(+) diff --git a/backend/open_webui/routers/users.py b/backend/open_webui/routers/users.py index 091143a5d5..3253ab3707 100644 --- a/backend/open_webui/routers/users.py +++ b/backend/open_webui/routers/users.py @@ -40,6 +40,7 @@ from open_webui.utils.auth import ( validate_password, ) from open_webui.utils.access_control import get_permissions, has_permission +from open_webui.socket.main import disconnect_user_sessions log = logging.getLogger(__name__) @@ -620,6 +621,10 @@ async def update_user_by_id( ) if updated_user: + # If the role changed, disconnect all socket sessions so stale + # privileges cached in SESSION_POOL are invalidated. + if updated_user.role != user.role: + await disconnect_user_sessions(user_id) return updated_user raise HTTPException( @@ -659,6 +664,7 @@ async def delete_user_by_id(user_id: str, user=Depends(get_admin_user), db: Asyn result = await Auths.delete_auth_by_id(user_id, db=db) if result: + await disconnect_user_sessions(user_id) return True raise HTTPException( diff --git a/backend/open_webui/socket/main.py b/backend/open_webui/socket/main.py index d2ddd90b19..2c44eb25c5 100644 --- a/backend/open_webui/socket/main.py +++ b/backend/open_webui/socket/main.py @@ -312,6 +312,24 @@ async def enter_room_for_users(room: str, user_ids: list[str]): log.debug(f'Failed to make users {user_ids} join room {room}: {e}') +async def disconnect_user_sessions(user_id: str): + """Disconnect all Socket.IO sessions belonging to a user. + + Call this when a user's role is changed or the user is deleted so that + stale role/permission data cached in SESSION_POOL is invalidated. + The client will automatically reconnect and re-authenticate with + fresh data from the database. + """ + try: + session_ids = get_session_ids_from_room(f'user:{user_id}') + for sid in session_ids: + await sio.disconnect(sid) + if session_ids: + log.info(f'Disconnected {len(session_ids)} session(s) for user {user_id}') + except Exception as e: + log.warning(f'Failed to disconnect sessions for user {user_id}: {e}') + + @sio.on('usage') async def usage(sid, data): if sid in SESSION_POOL: From fb5ef978bfb451c3f2221931e08e54250cec58ca Mon Sep 17 00:00:00 2001 From: Classic298 <27028174+Classic298@users.noreply.github.com> Date: Sun, 12 Apr 2026 23:19:58 +0200 Subject: [PATCH 36/67] fix: enforce OAUTH_ALLOWED_DOMAINS on token exchange endpoint (#23639) The OAuth token exchange endpoint skipped the domain allowlist check that the normal OAuth callback enforces. An attacker with a valid OAuth token from a non-allowed domain (e.g. gmail.com) could bypass the admin's domain restriction policy entirely. Adds the same domain validation check used in the OAuth callback, denying access when the email domain is not in the allowed list. --- backend/open_webui/routers/auths.py | 12 ++++++++++++ 1 file changed, 12 insertions(+) diff --git a/backend/open_webui/routers/auths.py b/backend/open_webui/routers/auths.py index 484212a493..fc3115dcbf 100644 --- a/backend/open_webui/routers/auths.py +++ b/backend/open_webui/routers/auths.py @@ -54,6 +54,7 @@ from open_webui.config import ( OAUTH_PROVIDERS, OAUTH_MERGE_ACCOUNTS_BY_EMAIL, ) +from open_webui.utils.oauth import auth_manager_config from pydantic import BaseModel from open_webui.utils.misc import parse_duration, validate_email_format @@ -1295,6 +1296,17 @@ async def token_exchange( ) email = email.lower() + # Enforce domain allowlist — same check as the normal OAuth callback + if ( + '*' not in auth_manager_config.OAUTH_ALLOWED_DOMAINS + and email.split('@')[-1] not in auth_manager_config.OAUTH_ALLOWED_DOMAINS + ): + log.warning(f'Token exchange denied: email domain not in allowed domains list') + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail=ERROR_MESSAGES.ACCESS_PROHIBITED, + ) + # Try to find the user by OAuth sub user = await Users.get_user_by_oauth_sub(provider, sub, db=db) From 5ee791d5d28f236755243cb7d16d8737bb69ce36 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Sun, 12 Apr 2026 16:25:01 -0500 Subject: [PATCH 37/67] refac --- backend/open_webui/env.py | 3 +++ backend/open_webui/main.py | 2 ++ backend/open_webui/utils/audit.py | 13 ++++++++++--- 3 files changed, 15 insertions(+), 3 deletions(-) diff --git a/backend/open_webui/env.py b/backend/open_webui/env.py index 9a16fd4ba8..1809268014 100644 --- a/backend/open_webui/env.py +++ b/backend/open_webui/env.py @@ -909,6 +909,9 @@ AUDIT_INCLUDED_PATHS = os.getenv('AUDIT_INCLUDED_PATHS', '').split(',') AUDIT_INCLUDED_PATHS = [path.strip() for path in AUDIT_INCLUDED_PATHS] AUDIT_INCLUDED_PATHS = [path.lstrip('/') for path in AUDIT_INCLUDED_PATHS if path] +# When enabled, GET requests are also audited (disabled by default to avoid log noise) +ENABLE_AUDIT_GET_REQUESTS = os.getenv('ENABLE_AUDIT_GET_REQUESTS', 'False').lower() == 'true' + #################################### # OPENTELEMETRY diff --git a/backend/open_webui/main.py b/backend/open_webui/main.py index d959351d4c..bee461bb2c 100644 --- a/backend/open_webui/main.py +++ b/backend/open_webui/main.py @@ -475,6 +475,7 @@ from open_webui.env import ( LICENSE_KEY, AUDIT_EXCLUDED_PATHS, AUDIT_INCLUDED_PATHS, + ENABLE_AUDIT_GET_REQUESTS, AUDIT_LOG_LEVEL, CHANGELOG, REDIS_URL, @@ -1560,6 +1561,7 @@ if audit_level != AuditLevel.NONE: audit_level=audit_level, excluded_paths=AUDIT_EXCLUDED_PATHS, included_paths=AUDIT_INCLUDED_PATHS, + audit_get_requests=ENABLE_AUDIT_GET_REQUESTS, max_body_size=MAX_BODY_LOG_SIZE, ) ################################## diff --git a/backend/open_webui/utils/audit.py b/backend/open_webui/utils/audit.py index 1200d813af..5686c88d5d 100644 --- a/backend/open_webui/utils/audit.py +++ b/backend/open_webui/utils/audit.py @@ -24,7 +24,7 @@ from asgiref.typing import ( from loguru import logger from starlette.requests import Request -from open_webui.env import AUDIT_LOG_LEVEL, AUDIT_INCLUDED_PATHS, MAX_BODY_LOG_SIZE +from open_webui.env import AUDIT_LOG_LEVEL, ENABLE_AUDIT_GET_REQUESTS, AUDIT_INCLUDED_PATHS, MAX_BODY_LOG_SIZE from open_webui.utils.auth import get_current_user, get_http_authorization_cred from open_webui.models.users import UserModel @@ -117,7 +117,7 @@ class AuditLoggingMiddleware: ASGI middleware that intercepts HTTP requests and responses to perform audit logging. It captures request/response bodies (depending on audit level), headers, HTTP methods, and user information, then logs a structured audit entry at the end of the request cycle. """ - AUDITED_METHODS = {'PUT', 'PATCH', 'DELETE', 'POST'} + DEFAULT_AUDITED_METHODS = {'PUT', 'PATCH', 'DELETE', 'POST'} def __init__( self, @@ -127,12 +127,16 @@ class AuditLoggingMiddleware: included_paths: Optional[list[str]] = None, max_body_size: int = MAX_BODY_LOG_SIZE, audit_level: AuditLevel = AuditLevel.NONE, + audit_get_requests: bool = False, ) -> None: self.app = app self.audit_logger = AuditLogger(logger) self.excluded_paths = excluded_paths or [] self.included_paths = included_paths or [] self.max_body_size = max_body_size + self.audited_methods = set(self.DEFAULT_AUDITED_METHODS) + if audit_get_requests: + self.audited_methods.add('GET') self.audit_level = audit_level if self.included_paths and self.excluded_paths: @@ -202,7 +206,10 @@ class AuditLoggingMiddleware: return None def _should_skip_auditing(self, request: Request) -> bool: - if request.method not in {'POST', 'PUT', 'PATCH', 'DELETE'} or AUDIT_LOG_LEVEL == 'NONE': + if AUDIT_LOG_LEVEL == 'NONE': + return True + + if request.method not in self.audited_methods: return True ALWAYS_LOG_ENDPOINTS = { From 4f94d21780909827ab2bf995a3169636175a2322 Mon Sep 17 00:00:00 2001 From: Classic298 <27028174+Classic298@users.noreply.github.com> Date: Sun, 12 Apr 2026 23:27:44 +0200 Subject: [PATCH 38/67] fix: enforce filter_allowed_access_grants on channel create and update (#23638) Unlike all other resource routers (knowledge, models, notes, prompts, tools, skills), the channel router did not call filter_allowed_access_grants. This allowed any user to set wildcard access grants on group channels, bypassing the admin's public sharing permission framework. Adds filter_allowed_access_grants with the sharing.public_channels permission key to both create and update endpoints, matching the pattern used by all other resource routers. --- backend/open_webui/routers/channels.py | 18 +++++++++++++++++- 1 file changed, 17 insertions(+), 1 deletion(-) diff --git a/backend/open_webui/routers/channels.py b/backend/open_webui/routers/channels.py index 6610ee2eca..5c2ab9dcac 100644 --- a/backend/open_webui/routers/channels.py +++ b/backend/open_webui/routers/channels.py @@ -61,7 +61,7 @@ from open_webui.utils.chat import generate_chat_completion from open_webui.utils.auth import get_admin_user, get_verified_user -from open_webui.utils.access_control import has_permission +from open_webui.utils.access_control import has_permission, filter_allowed_access_grants from open_webui.utils.webhook import post_webhook from open_webui.utils.channels import extract_mentions, replace_mentions from open_webui.internal.db import get_async_session @@ -303,6 +303,14 @@ async def create_new_channel( detail=ERROR_MESSAGES.UNAUTHORIZED, ) + form_data.access_grants = filter_allowed_access_grants( + request.app.state.config.USER_PERMISSIONS, + user.id, + user.role, + form_data.access_grants, + 'sharing.public_channels', + ) + try: if form_data.type == 'dm': existing_channel = await Channels.get_dm_channel_by_user_ids([user.id, *form_data.user_ids], db=db) @@ -633,6 +641,14 @@ async def update_channel_by_id( if channel.user_id != user.id and user.role != 'admin': raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) + form_data.access_grants = filter_allowed_access_grants( + request.app.state.config.USER_PERMISSIONS, + user.id, + user.role, + form_data.access_grants, + 'sharing.public_channels', + ) + try: channel = await Channels.update_channel_by_id(id, form_data, db=db) return ChannelModel(**channel.model_dump()) From 83024d00bbd56901ff7c4fa64590672e799ec416 Mon Sep 17 00:00:00 2001 From: Classic298 <27028174+Classic298@users.noreply.github.com> Date: Sun, 12 Apr 2026 23:33:41 +0200 Subject: [PATCH 39/67] fix: enforce API key endpoint restrictions at the auth layer, not middleware (#23637) The APIKeyRestrictionMiddleware only inspected the Authorization header for sk- tokens, but get_current_user also reads API keys from cookies and x-api-key headers. This allowed complete bypass of endpoint restrictions by sending the key via an alternate transport. Moves the restriction check into get_current_user_by_api_key so it runs regardless of how the API key was delivered. Removes the now-redundant middleware. --- backend/open_webui/main.py | 44 -------------------------------- backend/open_webui/utils/auth.py | 20 +++++++++++++++ 2 files changed, 20 insertions(+), 44 deletions(-) diff --git a/backend/open_webui/main.py b/backend/open_webui/main.py index bee461bb2c..620a8aa674 100644 --- a/backend/open_webui/main.py +++ b/backend/open_webui/main.py @@ -1391,50 +1391,6 @@ app.add_middleware(RedirectMiddleware) app.add_middleware(SecurityHeadersMiddleware) -class APIKeyRestrictionMiddleware: - def __init__(self, app): - self.app = app - - async def __call__(self, scope, receive, send): - if scope['type'] == 'http': - request = Request(scope) - auth_header = request.headers.get('Authorization') - token = None - - if auth_header: - parts = auth_header.split(' ', 1) - if len(parts) == 2: - token = parts[1] - - # Only apply restrictions if an sk- API key is used - if token and token.startswith('sk-'): - # Check if restrictions are enabled - if app.state.config.ENABLE_API_KEYS_ENDPOINT_RESTRICTIONS: - allowed_paths = [ - path.strip() - for path in str(app.state.config.API_KEYS_ALLOWED_ENDPOINTS).split(',') - if path.strip() - ] - - request_path = request.url.path - - # Match exact path or prefix path - is_allowed = any( - request_path == allowed or request_path.startswith(allowed + '/') for allowed in allowed_paths - ) - - if not is_allowed: - await JSONResponse( - status_code=status.HTTP_403_FORBIDDEN, - content={'detail': 'API key not allowed to access this endpoint.'}, - )(scope, receive, send) - return - - await self.app(scope, receive, send) - - -app.add_middleware(APIKeyRestrictionMiddleware) - @app.middleware('http') async def commit_session_after_request(request: Request, call_next): diff --git a/backend/open_webui/utils/auth.py b/backend/open_webui/utils/auth.py index 32e7db3423..9ac7524411 100644 --- a/backend/open_webui/utils/auth.py +++ b/backend/open_webui/utils/auth.py @@ -427,6 +427,26 @@ async def get_current_user_by_api_key(request, api_key: str): ): raise HTTPException(status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.API_KEY_NOT_ALLOWED) + # Enforce endpoint restrictions — checked here (not in middleware) + # so it applies regardless of how the API key was transported + # (Authorization header, cookie, x-api-key header, etc.). + if request.app.state.config.ENABLE_API_KEYS_ENDPOINT_RESTRICTIONS: + allowed_paths = [ + path.strip() + for path in str(request.app.state.config.API_KEYS_ALLOWED_ENDPOINTS).split(',') + if path.strip() + ] + request_path = request.url.path + is_allowed = any( + request_path == allowed or request_path.startswith(allowed + '/') + for allowed in allowed_paths + ) + if not is_allowed: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail=ERROR_MESSAGES.ACCESS_PROHIBITED, + ) + # Add user info to current span if ENABLE_OTEL: from opentelemetry import trace From b78dabb442dacd5430f0e8777aa65accf3d78c08 Mon Sep 17 00:00:00 2001 From: Classic298 <27028174+Classic298@users.noreply.github.com> Date: Sun, 12 Apr 2026 23:33:57 +0200 Subject: [PATCH 40/67] fix: reject empty passwords in LDAP authentication to prevent unauthenticated binds (#23633) Per RFC 4513, a Simple Bind with a non-empty DN but empty password is unauthenticated simple authentication. Many LDAP servers (OpenLDAP default, some AD configs) accept these binds, allowing account takeover without valid credentials. Rejects empty and whitespace-only passwords before attempting the LDAP bind. --- backend/open_webui/routers/auths.py | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/backend/open_webui/routers/auths.py b/backend/open_webui/routers/auths.py index fc3115dcbf..6652ebf44d 100644 --- a/backend/open_webui/routers/auths.py +++ b/backend/open_webui/routers/auths.py @@ -323,6 +323,14 @@ async def ldap_auth( detail=ERROR_MESSAGES.ACTION_PROHIBITED, ) + # Reject empty passwords before attempting the LDAP bind. + # Per RFC 4513 §5.1.2, a Simple Bind with a non-empty DN but empty + # password is "unauthenticated simple authentication" — many LDAP + # servers (OpenLDAP default, some AD configs) return success for these, + # which would grant access without valid credentials. + if not form_data.password or not form_data.password.strip(): + raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED) + # NOW load LDAP config variables LDAP_SERVER_LABEL = request.app.state.config.LDAP_SERVER_LABEL LDAP_SERVER_HOST = request.app.state.config.LDAP_SERVER_HOST From 0753409e7b18ffe4fc95a0a821f605d25b58ccad Mon Sep 17 00:00:00 2001 From: Classic298 <27028174+Classic298@users.noreply.github.com> Date: Sun, 12 Apr 2026 23:34:13 +0200 Subject: [PATCH 41/67] fix: use ipaddress stdlib for IPv6 SSRF protection (#23453) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The validators.ipv6(ip, private=True) call always returns a falsy ValidationError because validators==0.35.0 does not support the private kwarg for IPv6. This means any hostname resolving to a private IPv6 address (::1, fd00::*, ::ffff:169.254.169.254) bypasses SSRF protection entirely, circumventing the fix for CVE-2025-65958. Replace both the IPv4 and IPv6 validators-based private checks with Python's stdlib ipaddress module using an allowlist approach (not addr.is_global). This blocks all non-globally-routable addresses — private, loopback, link-local, reserved, multicast, and unspecified — for both IPv4 and IPv6, including IPv4-mapped IPv6 addresses. --- backend/open_webui/retrieval/web/utils.py | 9 ++++----- 1 file changed, 4 insertions(+), 5 deletions(-) diff --git a/backend/open_webui/retrieval/web/utils.py b/backend/open_webui/retrieval/web/utils.py index c9442f208b..cc520ffe63 100644 --- a/backend/open_webui/retrieval/web/utils.py +++ b/backend/open_webui/retrieval/web/utils.py @@ -1,4 +1,5 @@ import asyncio +import ipaddress import logging import socket import ssl @@ -84,11 +85,9 @@ def validate_url(url: Union[str, Sequence[str]]): ipv4_addresses, ipv6_addresses = resolve_hostname(parsed_url.hostname) # Check if any of the resolved addresses are private # This is technically still vulnerable to DNS rebinding attacks, as we don't control WebBaseLoader - for ip in ipv4_addresses: - if validators.ipv4(ip, private=True): - raise ValueError(ERROR_MESSAGES.INVALID_URL) - for ip in ipv6_addresses: - if validators.ipv6(ip, private=True): + for ip in ipv4_addresses + ipv6_addresses: + addr = ipaddress.ip_address(ip) + if not addr.is_global: raise ValueError(ERROR_MESSAGES.INVALID_URL) return True elif isinstance(url, Sequence): From 47d413ce7b2a006a8126f4a9055b13e5fcb33a1d Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Sun, 12 Apr 2026 16:47:23 -0500 Subject: [PATCH 42/67] refac --- backend/open_webui/routers/automations.py | 18 --------- backend/open_webui/utils/automations.py | 19 ++++++---- src/lib/components/AutomationModal.svelte | 36 +----------------- .../automations/AutomationEditor.svelte | 37 +------------------ src/lib/components/chat/Chat.svelte | 5 +++ .../workspace/Models/ModelEditor.svelte | 15 ++++++++ .../workspace/Models/TerminalSelector.svelte | 30 +++++++++++++++ 7 files changed, 65 insertions(+), 95 deletions(-) create mode 100644 src/lib/components/workspace/Models/TerminalSelector.svelte diff --git a/backend/open_webui/routers/automations.py b/backend/open_webui/routers/automations.py index 504bc726d2..9c532a8915 100644 --- a/backend/open_webui/routers/automations.py +++ b/backend/open_webui/routers/automations.py @@ -167,15 +167,6 @@ async def create_new_automation( await check_automation_limits(request, user, form_data.data.rrule, db, is_create=True) - # Validate terminal server exists if linked - if form_data.data.terminal and form_data.data.terminal.server_id: - connections = request.app.state.config.TERMINAL_SERVER_CONNECTIONS or [] - if not any(c.get('id') == form_data.data.terminal.server_id for c in connections): - raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail='Terminal server not found', - ) - tz = user.timezone automation = await Automations.insert(user.id, form_data, next_run_ns(form_data.data.rrule, tz=tz), db=db) return await enrich_automation(automation, db, tz=tz) @@ -226,15 +217,6 @@ async def update_automation_by_id( await check_automation_limits(request, user, form_data.data.rrule, db, is_create=False) - # Validate terminal server exists if linked - if form_data.data.terminal and form_data.data.terminal.server_id: - connections = request.app.state.config.TERMINAL_SERVER_CONNECTIONS or [] - if not any(c.get('id') == form_data.data.terminal.server_id for c in connections): - raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail='Terminal server not found', - ) - tz = user.timezone updated = await Automations.update_by_id(id, form_data, next_run_ns(form_data.data.rrule, tz=tz), db=db) return await enrich_automation(updated, db, tz=tz) diff --git a/backend/open_webui/utils/automations.py b/backend/open_webui/utils/automations.py index 262430dcf0..45b7ba65ab 100644 --- a/backend/open_webui/utils/automations.py +++ b/backend/open_webui/utils/automations.py @@ -224,6 +224,16 @@ def _resolve_model_filter_ids(app, model_id: str) -> list[str]: return list(filter_ids) if filter_ids else [] +def _resolve_model_terminal_id(app, model_id: str) -> Optional[str]: + """Read model default terminal_id from model config. + + The frontend does this in Chat.svelte (model.info.meta.terminalId). + """ + models = getattr(app.state, 'MODELS', {}) + model = models.get(model_id, {}) + return model.get('info', {}).get('meta', {}).get('terminalId') or None + + async def _set_terminal_cwd(app, server_id: str, user, cwd: str, chat_id: str) -> None: """Set the working directory on a terminal server via the proxy. @@ -357,13 +367,8 @@ async def execute_automation(app, automation: AutomationModel) -> None: features = _resolve_model_features(app, model_id) filter_ids = _resolve_model_filter_ids(app, model_id) - # If a terminal is linked, set the CWD before building the payload - terminal_id = None - if terminal_config and terminal_config.get('server_id'): - terminal_id = terminal_config['server_id'] - cwd = terminal_config.get('cwd') - if cwd: - await _set_terminal_cwd(app, terminal_id, user, cwd, chat.id) + # Resolve terminal from model config + terminal_id = _resolve_model_terminal_id(app, model_id) # Build the same payload the frontend sends to /api/chat/completions form_data = { diff --git a/src/lib/components/AutomationModal.svelte b/src/lib/components/AutomationModal.svelte index c1843f20e7..c16265515e 100644 --- a/src/lib/components/AutomationModal.svelte +++ b/src/lib/components/AutomationModal.svelte @@ -8,7 +8,6 @@ import ScheduleDropdown from '$lib/components/automations/ScheduleDropdown.svelte'; import ModelDropdown from '$lib/components/automations/ModelDropdown.svelte'; - import TerminalDropdown from '$lib/components/automations/TerminalDropdown.svelte'; import { createAutomation, @@ -16,7 +15,6 @@ type AutomationForm, type AutomationResponse } from '$lib/apis/automations'; - import { getTerminalServers, type TerminalServer } from '$lib/apis/terminal/index'; const i18n = getContext('i18n'); const dispatch = createEventDispatcher(); @@ -31,11 +29,6 @@ let loading = false; - // Terminal state - let terminalServers: TerminalServer[] = []; - let terminalServerId = ''; - let terminalCwd = ''; - // Schedule dropdown ref let scheduleDropdown: ScheduleDropdown; @@ -58,15 +51,7 @@ data: { prompt: prompt.trim(), model_id: model_id.trim(), - rrule: scheduleDropdown.buildRrule(), - ...(terminalServerId - ? { - terminal: { - server_id: terminalServerId, - ...(terminalCwd.trim() ? { cwd: terminalCwd.trim() } : {}) - } - } - : {}) + rrule: scheduleDropdown.buildRrule() }, is_active }; @@ -90,20 +75,11 @@ }; const init = async () => { - // Load terminal servers - try { - terminalServers = await getTerminalServers(localStorage.token); - } catch { - terminalServers = []; - } - if (automation) { name = automation.name; prompt = automation.data.prompt; model_id = automation.data.model_id; is_active = automation.is_active; - terminalServerId = automation.data.terminal?.server_id || ''; - terminalCwd = automation.data.terminal?.cwd || ''; if (scheduleDropdown) { scheduleDropdown.parseRrule(automation.data.rrule); } @@ -112,8 +88,6 @@ prompt = ''; model_id = ''; is_active = true; - terminalServerId = ''; - terminalCwd = ''; } }; @@ -158,14 +132,6 @@ - -
diff --git a/src/lib/components/automations/AutomationEditor.svelte b/src/lib/components/automations/AutomationEditor.svelte index fbdf448179..cb5859286a 100644 --- a/src/lib/components/automations/AutomationEditor.svelte +++ b/src/lib/components/automations/AutomationEditor.svelte @@ -19,7 +19,7 @@ type AutomationResponse, type AutomationRunModel } from '$lib/apis/automations'; - import { getTerminalServers, type TerminalServer } from '$lib/apis/terminal/index'; + import Spinner from '$lib/components/common/Spinner.svelte'; import Tooltip from '$lib/components/common/Tooltip.svelte'; @@ -29,7 +29,6 @@ import ScheduleDropdown from '$lib/components/automations/ScheduleDropdown.svelte'; import ModelDropdown from '$lib/components/automations/ModelDropdown.svelte'; - import TerminalDropdown from '$lib/components/automations/TerminalDropdown.svelte'; dayjs.extend(relativeTime); dayjs.extend(localizedFormat); @@ -43,9 +42,6 @@ let model_id = ''; let is_active = true; - let terminalServers: TerminalServer[] = []; - let terminalServerId = ''; - let terminalCwd = ''; let loading = false; let saving = false; @@ -97,15 +93,7 @@ data: { prompt: prompt.trim(), model_id: model_id.trim(), - rrule: scheduleDropdown.buildRrule(), - ...(terminalServerId - ? { - terminal: { - server_id: terminalServerId, - ...(terminalCwd.trim() ? { cwd: terminalCwd.trim() } : {}) - } - } - : {}) + rrule: scheduleDropdown.buildRrule() }, is_active }; @@ -204,18 +192,11 @@ prompt = automation.data.prompt; model_id = automation.data.model_id; is_active = automation.is_active; - terminalServerId = automation.data.terminal?.server_id || ''; - terminalCwd = automation.data.terminal?.cwd || ''; if (scheduleDropdown) { scheduleDropdown.parseRrule(automation.data.rrule); } - try { - terminalServers = await getTerminalServers(localStorage.token); - } catch { - terminalServers = []; - } await loadRuns(); }); @@ -355,20 +336,6 @@
- - {#if terminalServers.length > 0} -
- {$i18n.t('Terminal')} - -
- {/if} diff --git a/src/lib/components/chat/Chat.svelte b/src/lib/components/chat/Chat.svelte index a3feb6a269..6903429f7f 100644 --- a/src/lib/components/chat/Chat.svelte +++ b/src/lib/components/chat/Chat.svelte @@ -370,6 +370,11 @@ codeInterpreterEnabled = model.info.meta.defaultFeatureIds.includes('code_interpreter'); } } + + // Set Default Terminal + if (model?.info?.meta?.terminalId) { + selectedTerminalId.set(model.info.meta.terminalId); + } } }; diff --git a/src/lib/components/workspace/Models/ModelEditor.svelte b/src/lib/components/workspace/Models/ModelEditor.svelte index 5b5bd83caf..c41f51495a 100644 --- a/src/lib/components/workspace/Models/ModelEditor.svelte +++ b/src/lib/components/workspace/Models/ModelEditor.svelte @@ -25,6 +25,7 @@ import DefaultFeatures from './DefaultFeatures.svelte'; import BuiltinTools from './BuiltinTools.svelte'; import PromptSuggestions from './PromptSuggestions.svelte'; + import TerminalSelector from './TerminalSelector.svelte'; import AccessControlModal from '../common/AccessControlModal.svelte'; import LockClosed from '$lib/components/icons/LockClosed.svelte'; import { updateModelAccessGrants } from '$lib/apis/models'; @@ -102,6 +103,7 @@ let actionIds = []; let accessGrants = []; + let terminalId = ''; let tts = { voice: '' }; const submitHandler = async () => { @@ -206,6 +208,14 @@ } } + if (terminalId) { + info.meta.terminalId = terminalId; + } else { + if (info.meta.terminalId) { + delete info.meta.terminalId; + } + } + if (tts.voice !== '') { if (!info.meta.tts) info.meta.tts = {}; info.meta.tts.voice = tts.voice; @@ -316,6 +326,7 @@ capabilities = { ...capabilities, ...(model?.meta?.capabilities ?? {}) }; defaultFeatureIds = model?.meta?.defaultFeatureIds ?? defaultFeatureIds; builtinTools = model?.meta?.builtinTools ?? builtinTools; + terminalId = model?.meta?.terminalId ?? ''; tts = { voice: model?.meta?.tts?.voice ?? '' }; accessGrants = model?.access_grants ?? []; @@ -828,6 +839,10 @@ {/if} +
+ +
+
diff --git a/src/lib/components/workspace/Models/TerminalSelector.svelte b/src/lib/components/workspace/Models/TerminalSelector.svelte new file mode 100644 index 0000000000..501ec25d3f --- /dev/null +++ b/src/lib/components/workspace/Models/TerminalSelector.svelte @@ -0,0 +1,30 @@ + + +{#if terminals.length > 0} +
+
{$i18n.t('Terminal')}
+
+ + +{/if} From 008f1dfbdac5e539af051f0ea6be2b9d253ccc7c Mon Sep 17 00:00:00 2001 From: G30 <50341825+silentoplayz@users.noreply.github.com> Date: Sun, 12 Apr 2026 17:49:14 -0400 Subject: [PATCH 43/67] fix(ui): prevent user added action icons from being dragged (#23412) --- src/lib/components/chat/Messages/ResponseMessage.svelte | 1 + 1 file changed, 1 insertion(+) diff --git a/src/lib/components/chat/Messages/ResponseMessage.svelte b/src/lib/components/chat/Messages/ResponseMessage.svelte index 98b0d95a0b..35e592d139 100644 --- a/src/lib/components/chat/Messages/ResponseMessage.svelte +++ b/src/lib/components/chat/Messages/ResponseMessage.svelte @@ -1430,6 +1430,7 @@ : ''}" style="fill: currentColor;" alt={action.name} + draggable="false" />
{:else} From 15b89b9218b7d2c7239c579aa3d23c2892227ac6 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Sun, 12 Apr 2026 16:56:00 -0500 Subject: [PATCH 44/67] refac --- .../Markdown/MarkdownInlineTokens.svelte | 2 +- src/lib/utils/marked/katex-extension.ts | 51 +++++++++++++++++++ 2 files changed, 52 insertions(+), 1 deletion(-) diff --git a/src/lib/components/chat/Messages/Markdown/MarkdownInlineTokens.svelte b/src/lib/components/chat/Messages/Markdown/MarkdownInlineTokens.svelte index c2c7d81e61..0daf389eec 100644 --- a/src/lib/components/chat/Messages/Markdown/MarkdownInlineTokens.svelte +++ b/src/lib/components/chat/Messages/Markdown/MarkdownInlineTokens.svelte @@ -109,7 +109,7 @@ {:else if token.type === 'inlineKatex'} {#if token.text} - + {/if} {:else if token.type === 'iframe'}