diff --git a/backend/open_webui/config.py b/backend/open_webui/config.py
index e0e853e3ab..fe25dda1c1 100644
--- a/backend/open_webui/config.py
+++ b/backend/open_webui/config.py
@@ -915,6 +915,7 @@ if CUSTOM_NAME:
####################################
STORAGE_PROVIDER = os.environ.get('STORAGE_PROVIDER', 'local') # defaults to local, s3
+STORAGE_LOCAL_CACHE = os.environ.get('STORAGE_LOCAL_CACHE', 'true').lower() == 'true'
S3_ACCESS_KEY_ID = os.environ.get('S3_ACCESS_KEY_ID', None)
S3_SECRET_ACCESS_KEY = os.environ.get('S3_SECRET_ACCESS_KEY', None)
@@ -1541,6 +1542,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/env.py b/backend/open_webui/env.py
index 399cc3b603..8c1e8c2673 100644
--- a/backend/open_webui/env.py
+++ b/backend/open_webui/env.py
@@ -448,6 +448,27 @@ 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'
+
+# 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 == '':
@@ -518,6 +539,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)
####################################
@@ -800,6 +826,36 @@ else:
AIOHTTP_CLIENT_TIMEOUT_TOOL_SERVER = AIOHTTP_CLIENT_TIMEOUT
+####################################
+# AIOHTTP Connection Pool
+####################################
+
+AIOHTTP_POOL_CONNECTIONS = os.environ.get('AIOHTTP_POOL_CONNECTIONS', '')
+if AIOHTTP_POOL_CONNECTIONS == '':
+ AIOHTTP_POOL_CONNECTIONS = None
+else:
+ try:
+ AIOHTTP_POOL_CONNECTIONS = int(AIOHTTP_POOL_CONNECTIONS)
+ except ValueError:
+ AIOHTTP_POOL_CONNECTIONS = None
+
+AIOHTTP_POOL_CONNECTIONS_PER_HOST = os.environ.get('AIOHTTP_POOL_CONNECTIONS_PER_HOST', '')
+if AIOHTTP_POOL_CONNECTIONS_PER_HOST == '':
+ AIOHTTP_POOL_CONNECTIONS_PER_HOST = None
+else:
+ try:
+ AIOHTTP_POOL_CONNECTIONS_PER_HOST = int(AIOHTTP_POOL_CONNECTIONS_PER_HOST)
+ except ValueError:
+ AIOHTTP_POOL_CONNECTIONS_PER_HOST = None
+
+AIOHTTP_POOL_DNS_TTL = os.environ.get('AIOHTTP_POOL_DNS_TTL', '300')
+try:
+ AIOHTTP_POOL_DNS_TTL = int(AIOHTTP_POOL_DNS_TTL)
+ if AIOHTTP_POOL_DNS_TTL < 0:
+ AIOHTTP_POOL_DNS_TTL = 300
+except ValueError:
+ AIOHTTP_POOL_DNS_TTL = 300
+
RAG_EMBEDDING_TIMEOUT = os.environ.get('RAG_EMBEDDING_TIMEOUT', '')
if RAG_EMBEDDING_TIMEOUT == '':
@@ -903,6 +959,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/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 1279296000..6c30de2d55 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 pathlib import Path
from datetime import datetime, timedelta, timezone
@@ -30,6 +30,7 @@ from pydantic import BaseModel, SecretStr
from sqlalchemy import Dialect, create_engine, MetaData, event, types
from sqlalchemy.engine import Engine
from sqlalchemy.engine.url import URL
+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
@@ -222,6 +223,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')
@@ -299,6 +326,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)
@@ -306,6 +334,7 @@ ScopedSession = scoped_session(SessionLocal)
def get_session():
+ """Sync session generator — used ONLY for startup/config operations."""
db = SessionLocal()
try:
yield db
@@ -316,10 +345,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,
+ )
+ 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 03bb651089..8c2a139dba 100644
--- a/backend/open_webui/main.py
+++ b/backend/open_webui/main.py
@@ -67,6 +67,7 @@ from open_webui.socket.main import (
periodic_session_pool_cleanup,
get_event_emitter,
get_models_in_use,
+ get_user_id_from_session_pool,
)
from open_webui.routers import (
analytics,
@@ -108,8 +109,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
@@ -383,6 +384,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,
@@ -473,6 +476,7 @@ from open_webui.env import (
LICENSE_KEY,
AUDIT_EXCLUDED_PATHS,
AUDIT_INCLUDED_PATHS,
+ ENABLE_AUDIT_GET_REQUESTS,
AUDIT_LOG_LEVEL,
CHANGELOG,
REDIS_URL,
@@ -563,6 +567,7 @@ from open_webui.tasks import (
list_task_ids_by_item_id,
create_task,
stop_task,
+ stop_item_tasks,
list_tasks,
) # Import from tasks.py
@@ -573,7 +578,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__)
@@ -627,14 +632,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,
@@ -713,6 +721,11 @@ async def lifespan(app: FastAPI):
yield
+ # Shutdown: clean up shared resources
+ from open_webui.utils.session_pool import close_session
+
+ await close_session()
+
if hasattr(app.state, 'redis_task_command_listener'):
app.state.redis_task_command_listener.cancel()
@@ -874,6 +887,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
@@ -1383,51 +1398,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):
response = await call_next(request)
@@ -1553,6 +1523,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,
)
##################################
@@ -1601,7 +1572,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])}'
@@ -1667,12 +1638,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:
@@ -1754,7 +1725,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,
@@ -1766,7 +1737,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'),
[
@@ -1798,7 +1769,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'],
{
@@ -1809,13 +1780,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'},
@@ -1826,12 +1797,12 @@ async def chat_completion(
finally:
raise # re-raise to ensure proper task cancellation handling
except Exception as e:
- log.debug(f'Error processing chat payload: {e}')
+ log.error('Error processing chat payload: %s', e)
if metadata.get('chat_id') and metadata.get('message_id'):
# 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'],
{
@@ -1840,7 +1811,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',
@@ -1879,7 +1850,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:
@@ -1893,7 +1864,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}
@@ -2005,7 +1976,7 @@ async def chat_action(request: Request, action_id: str, form_data: dict, user=De
@app.post('/api/tasks/stop/{task_id}')
-async def stop_task_endpoint(request: Request, task_id: str, user=Depends(get_verified_user)):
+async def stop_task_endpoint(request: Request, task_id: str, user=Depends(get_admin_user)):
try:
result = await stop_task(request.app.state.redis, task_id)
return result
@@ -2014,15 +1985,21 @@ async def stop_task_endpoint(request: Request, task_id: str, user=Depends(get_ve
@app.get('/api/tasks')
-async def list_tasks_endpoint(request: Request, user=Depends(get_verified_user)):
+async def list_tasks_endpoint(request: Request, user=Depends(get_admin_user)):
return {'tasks': await list_tasks(request.app.state.redis)}
-@app.get('/api/tasks/chat/{chat_id}')
+@app.get('/api/tasks/chat/{chat_id:path}')
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)
- if chat is None or chat.user_id != user.id:
- return {'task_ids': []}
+ if chat_id.startswith('local:'):
+ socket_id = chat_id[len('local:') :]
+ owner_id = get_user_id_from_session_pool(socket_id)
+ if owner_id != user.id and user.role != 'admin':
+ return {'task_ids': []}
+ else:
+ chat = await Chats.get_chat_by_id(chat_id)
+ if chat is None or (chat.user_id != user.id and user.role != 'admin'):
+ return {'task_ids': []}
task_ids = await list_task_ids_by_item_id(request.app.state.redis, chat_id)
@@ -2030,6 +2007,21 @@ async def list_tasks_by_chat_id_endpoint(request: Request, chat_id: str, user=De
return {'task_ids': task_ids}
+@app.post('/api/tasks/chat/{chat_id:path}/stop')
+async def stop_tasks_by_chat_id_endpoint(request: Request, chat_id: str, user=Depends(get_verified_user)):
+ if chat_id.startswith('local:'):
+ socket_id = chat_id[len('local:') :]
+ owner_id = get_user_id_from_session_pool(socket_id)
+ if owner_id != user.id and user.role != 'admin':
+ raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND)
+ else:
+ chat = await Chats.get_chat_by_id(chat_id)
+ if chat is None or (chat.user_id != user.id and user.role != 'admin'):
+ raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND)
+ result = await stop_item_tasks(request.app.state.redis, chat_id)
+ return result
+
+
##################################
#
# Config Endpoints
@@ -2061,9 +2053,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:
@@ -2272,7 +2264,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
@@ -2479,7 +2471,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)
@@ -2492,7 +2484,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..f031495912 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,29 +282,28 @@ 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)
- .filter_by(
+ result = await db.execute(
+ select(AccessGrant).filter_by(
resource_type=resource_type,
resource_id=resource_id,
principal_type=principal_type,
principal_id=principal_id,
permission=permission,
)
- .first()
)
+ existing = result.scalars().first()
if existing:
return AccessGrantModel.model_validate(existing)
@@ -317,71 +317,69 @@ 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)
- .filter_by(
+ 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,
principal_type=principal_type,
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)
- .filter_by(
+ 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 +395,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 +433,77 @@ 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)
- .filter_by(
+ 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)
- .filter_by(
+ 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)
- .filter(
+ 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 +513,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 +532,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 +543,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 +573,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 +588,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 +599,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 +608,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 +626,20 @@ 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)
- .filter_by(
+ 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 +648,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 +670,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 +719,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 +777,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..2c8c6ba99f 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)
- .join(User, Auth.id == User.id)
- .filter(Auth.email == email, Auth.active == True)
- .first()
+ result = await db.execute(
+ select(Auth, User).join(User, Auth.id == User.id).filter(Auth.email == email, Auth.active == True)
)
- 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 485f097d5f..c891c3204e 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__)
@@ -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'),
+ )
####################
@@ -115,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()),
@@ -136,31 +139,36 @@ 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 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 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()
+
+ 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),
@@ -168,34 +176,37 @@ 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
@@ -205,37 +216,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.
@@ -243,7 +254,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(
@@ -257,7 +268,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
@@ -265,7 +277,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]
@@ -276,15 +288,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,
@@ -294,43 +306,71 @@ 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_by_automation(
+ 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 {}
+ async with get_async_db_context(db) as db:
+ # Subquery: max created_at per automation_id
+ subq = (
+ select(
+ AutomationRun.automation_id,
+ func.max(AutomationRun.created_at).label('max_created'),
+ )
+ .filter(AutomationRun.automation_id.in_(automation_ids))
+ .group_by(AutomationRun.automation_id)
+ .subquery()
+ )
+ result = await db.execute(
+ select(AutomationRun).join(
+ subq,
+ (AutomationRun.automation_id == subq.c.automation_id)
+ & (AutomationRun.created_at == subq.c.max_created),
+ )
+ )
+ rows = result.scalars().all()
+ return {row.automation_id: AutomationRunModel.model_validate(row) for row in rows}
+
+ 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 4d773491d5..942c06d6b3 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,22 @@ 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 +433,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 +442,32 @@ 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,58 +476,56 @@ 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)
@@ -548,19 +549,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
@@ -569,125 +569,131 @@ 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)
+ 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,
@@ -695,123 +701,127 @@ 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)
+ 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)
+ 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
@@ -823,16 +833,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()),
@@ -847,8 +857,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:
@@ -856,55 +866,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,
@@ -917,63 +930,70 @@ 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..bd9c720fa4 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,149 @@ 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(
+ 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(
+ 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,7 +359,7 @@ class ChatMessageTable:
else:
raise NotImplementedError(f'Unsupported dialect: {dialect}')
- query = db.query(
+ 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'),
@@ -366,14 +372,15 @@ class ChatMessageTable:
)
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 +389,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,7 +421,7 @@ class ChatMessageTable:
else:
raise NotImplementedError(f'Unsupported dialect: {dialect}')
- query = db.query(
+ 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'),
@@ -428,14 +434,15 @@ class ChatMessageTable:
)
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 +451,89 @@ 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(
+ 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(
+ 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(
+ 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 +555,29 @@ 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(
+ 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..b13edbf0bf 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,10 @@ 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 +313,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 +322,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 +350,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 +364,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 +373,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 +384,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 +440,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 +504,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 +513,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 +531,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 +550,21 @@ 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 +581,25 @@ 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 +607,102 @@ 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 +712,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 +739,25 @@ 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 +767,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 +795,47 @@ 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 +849,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 +857,32 @@ 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)
-
- if not include_folders:
- query = query.filter_by(folder_id=None)
-
- if not include_pinned:
- query = query.filter(or_(Chat.pinned == False, Chat.pinned == None))
-
- if not include_archived:
- query = query.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
+ 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:
+ stmt = stmt.filter_by(folder_id=None)
+
+ if not include_pinned:
+ stmt = stmt.filter(or_(Chat.pinned == False, Chat.pinned == None))
+
+ if not include_archived:
+ stmt = stmt.filter_by(archived=False)
+
+ 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 +896,104 @@ 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)
- .filter(Chat.id.in_(chat_ids))
- .filter_by(archived=False)
- .order_by(Chat.updated_at.desc())
- .all()
+ 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_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 +1001,25 @@ 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 +1028,16 @@ 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 +1051,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 +1073,19 @@ 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 +1123,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 +1158,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 +1175,7 @@ class ChatTable:
""")
)
elif tag_ids:
- query = query.filter(
+ stmt = stmt.filter(
and_(
*[
text(f"""
@@ -1183,14 +1191,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 +1208,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 +1225,7 @@ class ChatTable:
""")
)
elif tag_ids:
- query = query.filter(
+ stmt = stmt.filter(
and_(
*[
text(f"""
@@ -1236,39 +1240,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 +1289,82 @@ 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 +1378,54 @@ 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 +1437,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 +1469,138 @@ 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
- )
- db.query(ChatMessage).filter_by(chat_id=id).delete()
- db.query(Chat).filter_by(id=id).delete()
- db.commit()
+ async with get_async_db_context(db) as db:
+ await db.execute(update(AutomationRun).filter_by(chat_id=id).values(chat_id=None))
+ 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
- )
- db.query(ChatMessage).filter_by(chat_id=id).delete()
- db.query(Chat).filter_by(id=id, user_id=user_id).delete()
- db.commit()
+ async with get_async_db_context(db) as db:
+ await db.execute(update(AutomationRun).filter_by(chat_id=id).values(chat_id=None))
+ 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
- )
- db.query(Chat).filter_by(user_id=user_id, folder_id=folder_id).delete()
- db.commit()
+ await db.execute(delete(ChatMessage).filter(ChatMessage.chat_id.in_(chat_ids_stmt)))
+ 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 +1608,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 +1627,64 @@ 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)
- .filter_by(chat_id=chat_id, message_id=message_id)
- .order_by(ChatFile.created_at.asc())
- .all()
+ 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_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..02f61f82ee 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,99 @@ 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 +270,18 @@ 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 +289,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 +320,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 +375,22 @@ 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 +403,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 +428,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..cfdcfbc2d9 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,10 @@ 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 +148,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 +158,24 @@ 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 +183,14 @@ 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 +201,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,51 +215,53 @@ 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()
- 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()
- ]
+ 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 result.scalars().all()]
return FileListResponse(items=items, total=total)
@@ -275,13 +288,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 +309,26 @@ 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 +340,70 @@ 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..47dbe195ab 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,48 @@ 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)
- .filter_by(parent_id=parent_id, user_id=user_id)
- .filter(Folder.name.ilike(name))
- .first()
+ result = await db.execute(
+ select(Folder).filter_by(parent_id=parent_id, user_id=user_id).filter(Folder.name.ilike(name))
)
+ folder = result.scalars().first()
if not folder:
return None
@@ -182,25 +183,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 +208,38 @@ 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)
- .filter_by(
+ 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 +258,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 +279,41 @@ 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 +324,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 +335,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 +355,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..ddac317863 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,14 @@ 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 +181,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 +207,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 +217,29 @@ 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,36 @@ 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 +299,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 +335,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 +344,11 @@ 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 +362,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 +378,51 @@ 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..bc199fac5b 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,34 @@ 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 +274,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 +291,69 @@ 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,12 @@ 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 +383,104 @@ 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,33 @@ 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,15 +533,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)
-
- db.query(Group).filter(Group.id.in_(groups_to_remove)).update(
- {'updated_at': now}, synchronize_session=False
+ await db.execute(
+ delete(GroupMember).filter(
+ GroupMember.user_id == user_id,
+ GroupMember.group_id.in_(groups_to_remove),
+ )
)
+ await db.execute(update(Group).filter(Group.id.in_(groups_to_remove)).values(updated_at=now))
+
# 5. Bulk insert missing memberships
for group_id in groups_to_add:
db.add(
@@ -537,27 +555,26 @@ 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 +591,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 +606,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 +623,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..2750ef6058 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,10 +196,12 @@ class KnowledgeTable:
knowledge_bases.append(
KnowledgeUserModel.model_validate(
{
- **self._to_knowledge_model(
- knowledge,
- access_grants=grants_map.get(knowledge.id, []),
- db=db,
+ **(
+ await self._to_knowledge_model(
+ knowledge,
+ access_grants=grants_map.get(knowledge.id, []),
+ db=db,
+ )
).model_dump(),
'user': user.model_dump() if user else None,
}
@@ -206,22 +209,22 @@ class KnowledgeTable:
)
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,41 +236,45 @@ 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(
- knowledge_base,
- access_grants=grants_map.get(knowledge_base.id, []),
- db=db,
+ **(
+ await self._to_knowledge_model(
+ knowledge_base,
+ access_grants=grants_map.get(knowledge_base.id, []),
+ db=db,
+ )
).model_dump(),
'user': (UserModel.model_validate(user).model_dump() if user else None),
}
@@ -279,28 +286,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 +317,22 @@ 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,19 @@ 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 +515,36 @@ 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,107 @@ 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()),
- }
- )
- db.commit()
+ await db.execute(update(Knowledge).filter_by(id=id).values(updated_at=int(time.time())))
+ 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..7f33a72eff 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,39 @@ 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 +412,51 @@ 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..9bd3f888c1 100755
--- a/backend/open_webui/models/models.py
+++ b/backend/open_webui/models/models.py
@@ -1,9 +1,11 @@
+import json
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 +14,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 +153,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 +182,40 @@ 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,10 +223,12 @@ class ModelsTable:
models.append(
ModelUserResponse.model_validate(
{
- **self._to_model_model(
- model,
- access_grants=grants_map.get(model.id, []),
- db=db,
+ **(
+ await self._to_model_model(
+ model,
+ access_grants=grants_map.get(model.id, []),
+ db=db,
+ )
).model_dump(),
'user': user.model_dump() if user else None,
}
@@ -232,33 +236,38 @@ class ModelsTable:
)
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 +279,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,69 +306,82 @@ 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
- meta_text = func.lower(cast(Model.meta, String))
-
- query = query.filter(meta_text.like(like_pattern))
+ # SQLite stores JSON text via json.dumps(ensure_ascii=True),
+ # so non-ASCII chars are \uXXXX-escaped. PostgreSQL native JSONB
+ # stores literal Unicode. Use the right pattern for each.
+ if db.bind.dialect.name == 'sqlite':
+ if tag.isascii():
+ meta_text = func.lower(cast(Model.meta, String))
+ pattern = f'%{json.dumps(tag.lower())}%'
+ else:
+ meta_text = cast(Model.meta, String)
+ pattern = f'%{json.dumps(tag)}%'
+ else:
+ meta_text = func.lower(cast(Model.meta, String))
+ pattern = f'%{json.dumps(tag.lower(), ensure_ascii=False)}%'
+ stmt = stmt.filter(meta_text.like(pattern))
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(
- model,
- access_grants=grants_map.get(model.id, []),
- db=db,
+ **(
+ await self._to_model_model(
+ model,
+ access_grants=grants_map.get(model.id, []),
+ db=db,
+ )
).model_dump(),
user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
)
@@ -368,22 +389,23 @@ class ModelsTable:
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 +415,90 @@ 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 +507,14 @@ 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 +525,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..25a7905800 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,10 @@ 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,53 +131,58 @@ 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(
- or_(
- func.replace(func.replace(Note.title, '-', ''), ' ', '').ilike(f'%{normalized_query}%'),
- func.replace(
- func.replace(cast(Note.data['content']['md'], Text), '-', ''),
- ' ',
- '',
- ).ilike(f'%{normalized_query}%'),
+ # Split query into individual words and normalize each
+ # (strip hyphens so "todo" matches "to-do").
+ # All words must match somewhere in title OR content (AND semantics).
+ search_words = query_key.split()
+ normalized_words = [w.replace('-', '') for w in search_words if w.replace('-', '')]
+ for word in normalized_words:
+ stmt = stmt.filter(
+ or_(
+ func.replace(func.replace(Note.title, '-', ''), ' ', '').ilike(f'%{word}%'),
+ func.replace(
+ func.replace(cast(Note.data['content']['md'], Text), '-', ''),
+ ' ',
+ '',
+ ).ilike(f'%{word}%'),
+ )
)
- )
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 +190,9 @@ class NoteTable:
else:
permission = 'write'
- query = self._has_permission(
+ stmt = self._has_permission(
db,
- query,
+ stmt,
filter,
permission=permission,
)
@@ -195,46 +202,50 @@ 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(
- note,
- access_grants=grants_map.get(note.id, []),
- db=db,
+ **(
+ await self._to_note_model(
+ note,
+ access_grants=grants_map.get(note.id, []),
+ db=db,
+ )
).model_dump(),
user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
)
@@ -242,40 +253,44 @@ class NoteTable:
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 +304,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..84e5c66560 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,130 @@ 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 +261,71 @@ 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..fbf5401203 100644
--- a/backend/open_webui/models/prompts.py
+++ b/backend/open_webui/models/prompts.py
@@ -1,17 +1,19 @@
+import json
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 +94,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 +131,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 +150,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 +162,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,10 +214,12 @@ class PromptsTable:
prompts.append(
PromptUserResponse.model_validate(
{
- **self._to_prompt_model(
- prompt,
- access_grants=grants_map.get(prompt.id, []),
- db=db,
+ **(
+ await self._to_prompt_model(
+ prompt,
+ access_grants=grants_map.get(prompt.id, []),
+ db=db,
+ )
).model_dump(),
'user': user.model_dump() if user else None,
}
@@ -219,44 +228,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 +277,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',
@@ -284,55 +293,71 @@ class PromptsTable:
tag = filter.get('tag')
if tag:
- # 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))
+ # SQLite stores JSON text via json.dumps(ensure_ascii=True),
+ # so non-ASCII chars are \uXXXX-escaped. PostgreSQL native JSONB
+ # stores literal Unicode. Use the right pattern for each.
+ if db.bind.dialect.name == 'sqlite':
+ if tag.isascii():
+ tags_text = func.lower(cast(Prompt.tags, String))
+ pattern = f'%{json.dumps(tag.lower())}%'
+ else:
+ # LOWER() is ASCII-only; non-ASCII codepoints would
+ # produce different \uXXXX escapes when lowered.
+ tags_text = cast(Prompt.tags, String)
+ pattern = f'%{json.dumps(tag)}%'
+ else:
+ tags_text = func.lower(cast(Prompt.tags, String))
+ pattern = f'%{json.dumps(tag.lower(), ensure_ascii=False)}%'
+ stmt = stmt.filter(tags_text.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(
- prompt,
- access_grants=grants_map.get(prompt.id, []),
- db=db,
+ **(
+ await self._to_prompt_model(
+ prompt,
+ access_grants=grants_map.get(prompt.id, []),
+ db=db,
+ )
).model_dump(),
user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
)
@@ -340,22 +365,23 @@ class PromptsTable:
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 +397,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 +413,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 +425,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 +469,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 +488,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 +500,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 +529,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 +566,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..0fc6dfc52d 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,10 +184,12 @@ class SkillsTable:
skills.append(
SkillUserModel.model_validate(
{
- **self._to_skill_model(
- skill,
- access_grants=grants_map.get(skill.id, []),
- db=db,
+ **(
+ await self._to_skill_model(
+ skill,
+ access_grants=grants_map.get(skill.id, []),
+ db=db,
+ )
).model_dump(),
'user': user.model_dump() if user else None,
}
@@ -192,45 +197,45 @@ class SkillsTable:
)
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,43 +247,47 @@ 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(
- skill,
- access_grants=grants_map.get(skill.id, []),
- db=db,
+ **(
+ await self._to_skill_model(
+ skill,
+ access_grants=grants_map.get(skill.id, []),
+ db=db,
+ )
).model_dump(),
user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
)
@@ -289,43 +298,46 @@ 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..ee2baefc01 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,71 @@ 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..70035121aa 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,10 +172,12 @@ class ToolsTable:
tools.append(
ToolUserModel.model_validate(
{
- **self._to_tool_model(
- tool,
- access_grants=grants_map.get(tool.id, []),
- db=db,
+ **(
+ await self._to_tool_model(
+ tool,
+ access_grants=grants_map.get(tool.id, []),
+ db=db,
+ )
).model_dump(),
'user': user.model_dump() if user else None,
}
@@ -181,51 +185,57 @@ class ToolsTable:
)
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 +249,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 +265,34 @@ 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..025e79bd8a 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
@@ -247,20 +239,22 @@ class UserRoleUpdateForm(BaseModel):
class UserUpdateForm(BaseModel):
- role: str
- name: str
- email: str
- profile_image_url: str
+ role: Optional[str] = None
+ name: Optional[str] = None
+ email: Optional[str] = None
+ profile_image_url: Optional[str] = None
password: Optional[str] = None
- @field_validator('profile_image_url')
+ @field_validator('profile_image_url', mode='before')
@classmethod
- def check_profile_image_url(cls, v: str) -> str:
+ def check_profile_image_url(cls, v: Optional[str]) -> Optional[str]:
+ if v is None:
+ return v
return validate_profile_image_url(v)
class UsersTable:
- def insert_new_user(
+ async def insert_new_user(
self,
id: str,
name: str,
@@ -269,9 +263,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 +282,100 @@ 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 +384,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 +402,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 +420,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 +442,107 @@ 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))
- .join(GroupMember, User.id == GroupMember.user_id)
- .filter(GroupMember.group_id == group_id)
- .all()
+ 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)
)
+ 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 +551,75 @@ 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 +630,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 +643,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 +667,46 @@ 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 +717,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 +748,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 +772,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,9 +820,10 @@ 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
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/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):
diff --git a/backend/open_webui/routers/analytics.py b/backend/open_webui/routers/analytics.py
index 790c134295..fd045f79e7 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,12 @@ 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 +79,19 @@ 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 +122,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 +137,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 +156,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 +193,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 +228,12 @@ 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 +277,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 +297,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 +318,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 +363,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 +387,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 +431,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 9d8938b419..69744e7219 100644
--- a/backend/open_webui/routers/audio.py
+++ b/backend/open_webui/routers/audio.py
@@ -330,7 +330,9 @@ 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,
@@ -630,6 +632,7 @@ async def speech(request: Request, user=Depends(get_verified_user)):
detail=detail,
)
+
def transcription_handler(request, file_path, metadata, user=None):
filename = os.path.basename(file_path)
file_dir = os.path.dirname(file_path)
@@ -660,7 +663,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 +701,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 +770,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 +877,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 +1062,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)
@@ -1208,13 +1211,15 @@ 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,
@@ -1237,9 +1242,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/auths.py b/backend/open_webui/routers/auths.py
index 88f0fe69fb..cfc05160c1 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
@@ -70,8 +71,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 +97,9 @@ 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 +134,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 +170,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 +200,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 +230,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 +259,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 +281,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 +298,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 +313,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:
@@ -322,6 +325,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
@@ -476,23 +487,29 @@ 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:
- role = 'admin' if not Users.has_users(db=db) else request.app.state.config.DEFAULT_USER_ROLE
-
- user = Auths.insert_new_auth(
+ # Insert with default role first to avoid TOCTOU race on
+ # first-user registration. Matches signup_handler pattern.
+ user = await 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)
- apply_default_group_assignment(
+ # Atomically check if this is the only user *after* the
+ # insert. Only the single user present should become admin.
+ 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)
+
+ await apply_default_group_assignment(
request.app.state.config.DEFAULT_GROUP_ID,
user.id,
db=db,
@@ -504,19 +521,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:
@@ -536,7 +553,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(
@@ -558,7 +575,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,
@@ -567,20 +584,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}')
@@ -588,14 +605,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(
@@ -606,7 +623,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,
@@ -627,14 +644,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)
@@ -651,7 +668,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
@@ -665,7 +682,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,
@@ -678,9 +695,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:
@@ -695,7 +712,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,
@@ -709,9 +726,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:
@@ -724,7 +741,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:
@@ -741,7 +758,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:
@@ -750,7 +767,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')
@@ -771,7 +788,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.
@@ -846,12 +863,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:
@@ -861,7 +878,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,
@@ -871,13 +888,14 @@ 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,
)
- 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',
@@ -902,7 +920,9 @@ 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
@@ -910,11 +930,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
@@ -949,6 +969,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,
@@ -975,6 +997,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
@@ -1000,6 +1024,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
@@ -1041,6 +1071,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,
@@ -1154,10 +1186,12 @@ 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,
@@ -1165,7 +1199,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 {
@@ -1177,14 +1211,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,
@@ -1208,7 +1242,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.
@@ -1276,15 +1310,26 @@ 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 = 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(
@@ -1292,4 +1337,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 803f59a6f2..9c532a8915 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,
@@ -19,10 +19,11 @@ 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
-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__)
@@ -37,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(
@@ -60,8 +61,38 @@ def check_automation_access(automation, user):
)
-def enrich_automation(automation: AutomationModel, db: Session, tz: str = None) -> AutomationResponse:
- last_run = AutomationRuns.get_latest(automation.id, db=db)
+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
+
+ # 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 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})',
+ )
+
+ # 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.',
+ )
+
+
+async def enrich_automation(automation: AutomationModel, db: AsyncSession, tz: str = None) -> AutomationResponse:
+ """Full enrichment for single-item views (includes next_runs computation)."""
+ last_run = await AutomationRuns.get_latest(automation.id, db=db)
return AutomationResponse(
**automation.model_dump(),
last_run=last_run,
@@ -81,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,
@@ -97,8 +128,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 = await 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,
}
@@ -113,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:
@@ -124,18 +165,11 @@ async def create_new_automation(
detail=str(e),
)
- # 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',
- )
+ await check_automation_limits(request, user, form_data.data.rrule, db, is_create=True)
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)
############################
@@ -148,12 +182,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)
############################
@@ -167,10 +201,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:
@@ -181,18 +215,11 @@ async def update_automation_by_id(
detail=str(e),
)
- # 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',
- )
+ await check_automation_limits(request, user, form_data.data.rrule, db, is_create=False)
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)
############################
@@ -205,13 +232,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)
############################
@@ -224,13 +251,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)
############################
@@ -243,13 +270,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)
############################
@@ -264,9 +291,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 aa1ea52662..b6eee93eac 100644
--- a/backend/open_webui/routers/channels.py
+++ b/backend/open_webui/routers/channels.py
@@ -61,25 +61,25 @@ 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_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,10 @@ 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 +135,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 +144,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 +170,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 +190,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 +218,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 +236,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 +253,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 +267,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 +294,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)
@@ -303,12 +305,20 @@ 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 = 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 +327,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 +368,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 +379,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 +404,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 +412,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 +423,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 +435,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 +461,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,16 +475,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 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:
@@ -496,13 +509,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,
}
@@ -522,17 +535,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
@@ -552,10 +565,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)
@@ -563,7 +576,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
)
@@ -588,11 +601,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)
@@ -600,7 +613,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:
@@ -619,19 +632,27 @@ 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)
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 = 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)
@@ -648,11 +669,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)
@@ -660,7 +681,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)
@@ -692,40 +713,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(
@@ -733,7 +754,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,
}
)
@@ -755,32 +776,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:
@@ -788,12 +809,12 @@ 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:
- 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
@@ -801,7 +822,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,
}
)
@@ -820,12 +841,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,
@@ -843,7 +864,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)
@@ -874,10 +895,12 @@ 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(
- channel.id,
- message.parent_id if message.parent_id else message.id,
- db=db,
+ thread_messages = (
+ await Messages.get_messages_by_parent_id(
+ channel.id,
+ message.parent_id if message.parent_id else message.id,
+ db=db,
+ )
)[::-1]
response_message, channel = await new_message_handler(
@@ -905,7 +928,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]
@@ -925,7 +948,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)
@@ -1014,15 +1037,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',
@@ -1032,15 +1055,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,
@@ -1060,7 +1083,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(
@@ -1092,16 +1115,18 @@ 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)
@@ -1141,31 +1166,32 @@ 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 = await 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,
}
)
@@ -1181,21 +1207,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)
@@ -1221,21 +1247,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)
@@ -1243,12 +1269,13 @@ 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)
+ 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(),
- '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:
@@ -1269,35 +1296,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(
@@ -1305,7 +1332,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,
}
)
@@ -1326,14 +1353,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)
@@ -1341,19 +1368,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(
@@ -1393,18 +1420,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',
@@ -1413,7 +1440,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)
@@ -1421,8 +1448,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',
@@ -1460,18 +1487,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',
@@ -1480,7 +1507,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)
@@ -1488,9 +1515,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',
@@ -1527,14 +1554,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)
@@ -1542,13 +1569,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',
@@ -1559,7 +1586,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',
{
@@ -1580,7 +1607,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(
@@ -1610,9 +1637,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')
@@ -1648,18 +1675,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)
@@ -1668,18 +1695,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())
@@ -1693,22 +1720,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())
@@ -1721,22 +1748,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)
############################
@@ -1754,25 +1781,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
@@ -1786,10 +1813,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..979b9388cf 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,17 @@ 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 +518,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 +540,9 @@ 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 +554,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 +573,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 +589,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 +603,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 +611,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 +624,17 @@ 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 +643,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 +666,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 +676,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 +687,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 +697,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 +712,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 +733,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 +749,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 +764,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 +774,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 +790,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 +806,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 +821,16 @@ 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 +857,13 @@ 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 +874,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 +894,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 +921,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 +937,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 +945,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 +983,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 +999,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 +1027,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 +1066,10 @@ 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 +1082,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 +1105,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 +1116,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 +1149,13 @@ 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 +1165,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 +1198,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 +1226,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 +1264,16 @@ 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 +1297,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 +1313,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 +1332,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 +1346,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 +1365,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 +1387,14 @@ 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..072c7fa732 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,13 @@ 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 48172e744e..66d7539278 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, get_async_db_context
from open_webui.constants import ERROR_MESSAGES
from open_webui.retrieval.vector.factory import VECTOR_DB_CLIENT
@@ -48,7 +48,7 @@ from open_webui.routers.audio import transcribe
from open_webui.storage.provider import Storage
-from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL
+from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL, STORAGE_LOCAL_CACHE, STORAGE_PROVIDER, UPLOAD_DIR
from open_webui.utils.auth import get_admin_user, get_verified_user
from open_webui.utils.misc import strict_match_mime_type
from pydantic import BaseModel
@@ -88,16 +88,30 @@ def _is_text_file(file_path: str, chunk_size: int = 8192) -> bool:
return False
-def process_uploaded_file(
+def _cleanup_local_cache(file_path: str) -> None:
+ """Remove the local cached copy of a cloud-stored file after processing."""
+ if STORAGE_LOCAL_CACHE or STORAGE_PROVIDER == 'local':
+ return
+ try:
+ local_filename = os.path.basename(file_path)
+ local_path = os.path.join(UPLOAD_DIR, local_filename)
+ if os.path.isfile(local_path):
+ os.remove(local_path)
+ log.debug(f'Cleaned up local cache: {local_path}')
+ except OSError as e:
+ log.warning(f'Failed to clean up local cache for {file_path}: {e}')
+
+
+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
@@ -113,7 +127,7 @@ 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 +136,7 @@ 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 +146,7 @@ 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,
@@ -141,7 +155,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',
@@ -150,15 +164,18 @@ def process_uploaded_file(
db=db_session,
)
- if db:
- _process_handler(db)
- else:
- with SessionLocal() as db_session:
- _process_handler(db_session)
+ try:
+ if db:
+ await _process_handler(db)
+ else:
+ async with get_async_db_context() as db_session:
+ await _process_handler(db_session)
+ finally:
+ _cleanup_local_cache(file_path)
@router.post('/', response_model=FileModelResponse)
-def upload_file(
+async def upload_file(
request: Request,
background_tasks: BackgroundTasks,
file: UploadFile = File(...),
@@ -166,9 +183,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 +197,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 +205,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 +253,7 @@ def upload_file_handler(
},
)
- file_item = Files.insert_new_file(
+ file_item = await Files.insert_new_file(
user.id,
FileForm(
**{
@@ -258,9 +275,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 +292,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 +334,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 +364,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 +374,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 +402,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 +429,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 +438,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 +452,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 +462,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 +471,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 +512,10 @@ 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 +523,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 +542,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,15 +557,15 @@ 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(
+ await process_file(
request,
ProcessFileForm(file_id=id, content=form_data.content),
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,13 +573,13 @@ 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
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,
@@ -587,9 +606,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 +616,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 +665,10 @@ 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 +676,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)
- if not file_user.role == 'admin':
+ 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 +714,10 @@ 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 +725,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 +772,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 +781,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 +795,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..ebd0c0cb17 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,31 @@ 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 +104,14 @@ 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 +120,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 +137,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 +158,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 +174,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 +204,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 +219,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 +249,14 @@ 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 +283,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 +296,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 +319,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..de09aa05a1 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,13 @@ 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 +403,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 +434,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 +448,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 +475,13 @@ 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 +500,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 +526,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 +540,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 dca9a58a7a..0d534db7f6 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,
@@ -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,
)
@@ -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,
@@ -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,
@@ -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,
)
@@ -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..c6f8ce5ecd 100644
--- a/backend/open_webui/routers/knowledge.py
+++ b/backend/open_webui/routers/knowledge.py
@@ -2,14 +2,14 @@ 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
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)
@@ -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,
@@ -357,12 +356,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 +384,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 +404,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 +437,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 +450,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 +463,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 +471,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 +482,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 +506,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 +517,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 +531,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 +539,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 +561,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 +573,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 +601,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 +614,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 +630,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 +644,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,
@@ -659,7 +658,7 @@ 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,
@@ -667,7 +666,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 +677,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 +687,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 +703,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 +717,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 +725,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,
@@ -737,7 +736,7 @@ 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,
@@ -752,7 +751,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 +766,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 +782,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 +796,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 +804,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 +838,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 +858,10 @@ 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 +870,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 +887,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 +911,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 +923,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 +933,10 @@ 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 +945,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,
@@ -962,7 +965,7 @@ async def reset_knowledge_by_id(id: str, user=Depends(get_verified_user), db: Se
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
@@ -977,12 +980,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 +994,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 +1011,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 +1037,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 +1053,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 +1063,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 93745440c4..0c6fe73abf 100644
--- a/backend/open_webui/routers/ollama.py
+++ b/backend/open_webui/routers/ollama.py
@@ -39,17 +39,21 @@ 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
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,
+)
+from open_webui.utils.session_pool import (
cleanup_response,
+ get_session,
stream_wrapper,
)
from open_webui.utils.payload import (
@@ -121,10 +125,7 @@ async def send_request(
r = None
streaming = False
try:
- session = aiohttp.ClientSession(
- trust_env=True,
- timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT),
- )
+ session = await get_session()
headers = {
'Content-Type': 'application/json',
@@ -137,8 +138,12 @@ async def send_request(
headers[FORWARD_SESSION_INFO_HEADER_CHAT_ID] = metadata.get('chat_id')
r = await session.request(
- method, url, data=payload, headers=headers,
+ method,
+ url,
+ data=payload,
+ headers=headers,
ssl=AIOHTTP_CLIENT_SESSION_SSL,
+ timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT),
)
if not r.ok:
@@ -164,7 +169,7 @@ async def send_request(
streaming = True
return StreamingResponse(
- stream_wrapper(r, session),
+ stream_wrapper(r),
status_code=r.status,
headers=response_headers,
)
@@ -183,7 +188,7 @@ async def send_request(
)
finally:
if not streaming:
- await cleanup_response(r, session)
+ await cleanup_response(r)
def get_api_key(idx, url, configs):
@@ -397,11 +402,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()),
@@ -780,7 +785,8 @@ async def delete_model(
key = get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS)
await send_request(
- f'{url}/api/delete', 'DELETE',
+ f'{url}/api/delete',
+ 'DELETE',
payload=json.dumps(form_data),
key=key,
user=user,
@@ -796,11 +802,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
+ 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
- model = form_data.get('model')
-
if model not in models:
raise HTTPException(
status_code=400,
@@ -845,6 +854,9 @@ async def embed(
log.info(f'generate_ollama_batch_embeddings {form_data}')
+ # Enforce per-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
@@ -901,6 +913,9 @@ async def embeddings(
log.info(f'generate_ollama_embeddings {form_data}')
+ # Enforce per-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
@@ -963,11 +978,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
+ 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)
models = request.app.state.OLLAMA_MODELS
model = form_data.model
+
if model in models:
url_idx = random.choice(models[model]['urls'])
else:
@@ -1050,7 +1069,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.
@@ -1079,7 +1098,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:
@@ -1097,29 +1116,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',
- )
+ await check_model_access(user, model_info, bypass_filter)
+ else:
+ 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(
@@ -1177,7 +1176,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.
@@ -1197,7 +1196,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
@@ -1206,29 +1205,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',
- )
+ await check_model_access(user, model_info)
else:
- if user.role != 'admin':
- raise HTTPException(
- status_code=403,
- detail='Model not found',
- )
+ 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(
@@ -1259,7 +1238,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.
@@ -1279,7 +1258,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
@@ -1292,29 +1271,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',
- )
+ await check_model_access(user, model_info)
else:
- if user.role != 'admin':
- raise HTTPException(
- status_code=403,
- detail='Model not found',
- )
+ 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(
@@ -1359,34 +1318,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 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',
- )
+ await check_model_access(user, model_info)
else:
- if user.role != 'admin':
- raise HTTPException(
- status_code=403,
- detail='Model not found',
- )
+ 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(
@@ -1437,17 +1376,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,
@@ -1492,7 +1431,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:
@@ -1524,11 +1463,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()),
@@ -1651,7 +1590,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),
@@ -1715,9 +1654,7 @@ async def upload_model(
url = f'{ollama_url}/api/blobs/sha256:{file_hash}'
upload_timeout = aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT)
async with aiohttp.ClientSession(timeout=upload_timeout, trust_env=True) as upload_session:
- async with upload_session.post(
- url, data=blob_data, ssl=AIOHTTP_CLIENT_SESSION_SSL
- ) as response:
+ async with upload_session.post(url, data=blob_data, ssl=AIOHTTP_CLIENT_SESSION_SSL) as response:
if not response.ok:
raise Exception('Ollama: Could not create blob, Please try again.')
diff --git a/backend/open_webui/routers/openai.py b/backend/open_webui/routers/openai.py
index 836517df9d..09db29e6d6 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
@@ -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,
@@ -21,13 +21,14 @@ 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
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,
)
@@ -39,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
@@ -50,9 +52,12 @@ from open_webui.utils.payload import (
apply_system_prompt_to_body,
)
from open_webui.utils.misc import (
- cleanup_response,
convert_logit_bias_input_to_json,
stream_chunks_handler,
+)
+from open_webui.utils.session_pool import (
+ cleanup_response,
+ get_session,
stream_wrapper,
)
@@ -449,11 +454,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()),
@@ -771,6 +776,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', '')
@@ -794,6 +814,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
@@ -1006,7 +1029,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.
@@ -1024,7 +1047,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:
@@ -1044,29 +1067,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',
- )
+ await check_model_access(user, model_info, bypass_filter)
+ else:
+ 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
@@ -1174,12 +1177,11 @@ async def generate_chat_completion(
payload = json.dumps(payload)
r = None
- session = None
streaming = False
response = None
try:
- session = aiohttp.ClientSession(trust_env=True, timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT))
+ session = await get_session()
r = await session.request(
method='POST',
@@ -1188,13 +1190,33 @@ async def generate_chat_completion(
headers=headers,
cookies=cookies,
ssl=AIOHTTP_CLIENT_SESSION_SSL,
+ timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT),
)
# Check if response is SSE
if 'text/event-stream' in r.headers.get('Content-Type', ''):
+ # If the provider returned an error status with SSE content-type,
+ # read the body and return a proper error response instead of
+ # streaming the error back (which hides the error from logs).
+ if r.status >= 400:
+ error_body = await r.text()
+ log.error(
+ 'Provider returned HTTP %d with SSE content-type: %s',
+ r.status,
+ error_body[:1000],
+ )
+ try:
+ error_json = json.loads(error_body)
+ return JSONResponse(status_code=r.status, content=error_json)
+ except json.JSONDecodeError:
+ return JSONResponse(
+ status_code=r.status,
+ content={'error': {'message': error_body, 'code': r.status}},
+ )
+
streaming = True
return StreamingResponse(
- stream_wrapper(r, session, stream_chunks_handler),
+ stream_wrapper(r, content_handler=stream_chunks_handler),
status_code=r.status,
headers=dict(r.headers),
)
@@ -1225,7 +1247,7 @@ async def generate_chat_completion(
)
finally:
if not streaming:
- await cleanup_response(r, session)
+ await cleanup_response(r)
async def embeddings(request: Request, form_data: dict, user):
@@ -1261,27 +1283,24 @@ async def embeddings(request: Request, form_data: dict, user):
)
r = None
- session = None
streaming = False
headers, cookies = await get_headers_and_cookies(request, url, key, api_config, user=user)
try:
- session = aiohttp.ClientSession(
- trust_env=True,
- timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT),
- )
+ session = await get_session()
r = await session.request(
method='POST',
url=f'{url}/embeddings',
data=body,
headers=headers,
cookies=cookies,
+ timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT),
)
if 'text/event-stream' in r.headers.get('Content-Type', ''):
streaming = True
return StreamingResponse(
- stream_wrapper(r, session),
+ stream_wrapper(r),
status_code=r.status,
headers=dict(r.headers),
)
@@ -1306,7 +1325,7 @@ async def embeddings(request: Request, form_data: dict, user):
)
finally:
if not streaming:
- await cleanup_response(r, session)
+ await cleanup_response(r)
class ResponsesForm(BaseModel):
@@ -1340,10 +1359,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
+ await check_model_access(user, await 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:
@@ -1360,7 +1384,6 @@ async def responses(
)
r = None
- session = None
streaming = False
try:
@@ -1378,15 +1401,12 @@ 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'
- session = aiohttp.ClientSession(
- trust_env=True,
- timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT),
- )
+ session = await get_session()
r = await session.request(
method='POST',
url=request_url,
@@ -1394,13 +1414,14 @@ async def responses(
headers=headers,
cookies=cookies,
ssl=AIOHTTP_CLIENT_SESSION_SSL,
+ timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT),
)
# Check if response is SSE
if 'text/event-stream' in r.headers.get('Content-Type', ''):
streaming = True
return StreamingResponse(
- stream_wrapper(r, session),
+ stream_wrapper(r),
status_code=r.status,
headers=dict(r.headers),
)
@@ -1418,6 +1439,8 @@ async def responses(
return response_data
+ except HTTPException:
+ raise
except Exception as e:
log.exception(e)
raise HTTPException(
@@ -1426,15 +1449,22 @@ async def responses(
)
finally:
if not streaming:
- await cleanup_response(r, session)
+ await cleanup_response(r)
@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
@@ -1465,7 +1495,6 @@ async def proxy(path: str, request: Request, user=Depends(get_verified_user)):
)
r = None
- session = None
streaming = False
try:
@@ -1494,10 +1523,7 @@ async def proxy(path: str, request: Request, user=Depends(get_verified_user)):
else:
request_url = f'{url}/{path}'
- session = aiohttp.ClientSession(
- trust_env=True,
- timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT),
- )
+ session = await get_session()
r = await session.request(
method=request.method,
url=request_url,
@@ -1505,13 +1531,14 @@ async def proxy(path: str, request: Request, user=Depends(get_verified_user)):
headers=headers,
cookies=cookies,
ssl=AIOHTTP_CLIENT_SESSION_SSL,
+ timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT),
)
# Check if response is SSE
if 'text/event-stream' in r.headers.get('Content-Type', ''):
streaming = True
return StreamingResponse(
- stream_wrapper(r, session),
+ stream_wrapper(r),
status_code=r.status,
headers=dict(r.headers),
)
@@ -1529,6 +1556,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(
@@ -1537,4 +1566,4 @@ async def proxy(path: str, request: Request, user=Depends(get_verified_user)):
)
finally:
if not streaming:
- await cleanup_response(r, session)
+ await cleanup_response(r)
diff --git a/backend/open_webui/routers/prompts.py b/backend/open_webui/routers/prompts.py
index 3b579c2892..a4a75754f0 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,16 @@ 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 +220,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 +242,16 @@ 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 +264,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 +291,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 +304,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 +320,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 +336,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 +356,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 +369,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 +385,16 @@ 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 +409,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 +420,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 +434,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 +459,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 +470,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 +484,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 +492,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 +503,10 @@ 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 +516,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 +530,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 +545,10 @@ 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 +558,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 +572,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 +586,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 +603,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 +616,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 +625,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 +640,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 +653,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 +668,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 +683,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 +703,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 +719,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 +734,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 +747,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..1e1beac4fa 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_db, get_async_session
+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,
@@ -1694,10 +1694,10 @@ def process_file(
try:
# Commit any pending changes before the slow embedding step.
# Note: file is already a Pydantic model (not ORM), so no expunge needed.
- db.commit()
+ await db.commit()
# External embedding API takes time (5-60s+).
- # Subsequent updates use fresh sessions via get_db().
+ # Subsequent updates use fresh async sessions.
result = save_docs_to_vector_db(
request,
docs=docs,
@@ -1714,8 +1714,8 @@ def process_file(
if result:
# Fresh session for the final update.
- with get_db() as session:
- Files.update_file_metadata_by_id(
+ async with get_async_db() as session:
+ 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,
@@ -1744,14 +1744,14 @@ def process_file(
except Exception as e:
log.exception(e)
# Fresh session for error status update.
- with get_db() as session:
- Files.update_file_data_by_id(
+ async with get_async_db() as session:
+ 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 56923bc447..75f45bcaf9 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
@@ -29,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__)
@@ -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',
@@ -325,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 ['', '']
@@ -344,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,
@@ -378,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,
@@ -511,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):
@@ -526,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,
@@ -559,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)
@@ -574,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,
@@ -595,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,
@@ -618,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,
@@ -636,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)
@@ -648,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,
@@ -681,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,
@@ -691,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)
@@ -703,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,
@@ -733,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,
@@ -746,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)
@@ -754,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,
@@ -782,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):
@@ -794,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)
@@ -809,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,
@@ -824,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)
@@ -842,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
@@ -860,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,
@@ -883,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)
@@ -897,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,
@@ -918,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)
@@ -937,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,
@@ -964,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':
@@ -972,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)
@@ -995,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..041a866e37 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,31 +159,30 @@ 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)}
- tools = [
- tool
- for tool in tools
- if tool.user_id == user.id
- or (
- has_access(
+ user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)}
+ filtered_tools = []
+ for tool in tools:
+ 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 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
############################
@@ -192,13 +191,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 +297,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 +311,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 +326,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 +350,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 +362,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 +400,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 +420,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 +452,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 +464,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 +480,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 +488,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 +502,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 +534,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 +545,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 +559,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 +567,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 +582,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 +593,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 +607,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 +622,10 @@ 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 +634,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 +649,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 +668,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 +679,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 +696,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 +719,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 +730,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 +747,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 +761,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 +777,10 @@ 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 +789,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 +804,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 +818,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 +829,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 +846,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 +864,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 +875,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 +892,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 +902,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..9fd2479ada 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 (
@@ -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__)
@@ -63,7 +64,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 +81,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 +107,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 +119,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 +134,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 +143,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 +156,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 +273,10 @@ 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 +296,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 +304,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 +313,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 +332,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 +359,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 +383,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 +401,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 +438,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 +452,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 +470,17 @@ 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 +491,10 @@ 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 +510,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 +549,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 +566,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:
@@ -573,7 +580,7 @@ async def update_user_by_id(
detail=ERROR_MESSAGES.ACTION_PROHIBITED,
)
- if form_data.role != 'admin':
+ if form_data.role is not None and form_data.role != 'admin':
# If the primary admin is trying to change their own role, prevent it
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
@@ -587,11 +594,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)
+ if form_data.email is not None and form_data.email.lower() != user.email:
+ 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,21 +612,34 @@ 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(
- user_id,
- {
- 'role': form_data.role,
- 'name': form_data.name,
- 'email': form_data.email.lower(),
- 'profile_image_url': form_data.profile_image_url,
- },
- db=db,
- )
+ # Build update dict from only the provided fields
+ update_data = {}
+ if form_data.role is not None:
+ update_data['role'] = form_data.role
+ if form_data.name is not None:
+ update_data['name'] = form_data.name
+ if form_data.email is not None:
+ update_data['email'] = form_data.email.lower()
+ await Auths.update_email_by_id(user_id, form_data.email.lower(), db=db)
+ if form_data.profile_image_url is not None:
+ update_data['profile_image_url'] = form_data.profile_image_url
+
+ if update_data:
+ updated_user = await Users.update_user_by_id(
+ user_id,
+ update_data,
+ db=db,
+ )
+ else:
+ updated_user = user
if updated_user:
+ # If the role changed, disconnect all socket sessions so stale
+ # privileges cached in SESSION_POOL are invalidated.
+ if updated_user.role != user.role:
+ await disconnect_user_sessions(user_id)
return updated_user
raise HTTPException(
@@ -639,10 +659,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,9 +676,10 @@ 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:
+ await disconnect_user_sessions(user_id)
return True
raise HTTPException(
@@ -679,5 +700,7 @@ 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..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:
@@ -333,7 +351,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 +379,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 +399,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 +413,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 +426,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 +448,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 +460,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 +506,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 +519,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 +547,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 +555,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 +620,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 +628,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 +638,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 +811,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 +831,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 +847,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 +858,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 +867,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 +875,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 +884,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 +892,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 +903,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 +911,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 +925,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/storage/provider.py b/backend/open_webui/storage/provider.py
index 3c29462349..f70f3e862b 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:
@@ -140,7 +140,7 @@ class S3StorageProvider(StorageProvider):
def upload_file(self, file: BinaryIO, filename: str, tags: Dict[str, str]) -> Tuple[bytes, str]:
"""Handles uploading of the file to S3 storage."""
- _, file_path = LocalStorageProvider.upload_file(file, filename, tags)
+ contents, file_path = LocalStorageProvider.upload_file(file, filename, tags)
s3_key = os.path.join(self.key_prefix, filename)
try:
self.s3_client.upload_file(file_path, self.bucket_name, s3_key)
@@ -153,7 +153,7 @@ class S3StorageProvider(StorageProvider):
Tagging=tagging,
)
return (
- open(file_path, 'rb').read(),
+ contents,
f's3://{self.bucket_name}/{s3_key}',
)
except ClientError as e:
@@ -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())
diff --git a/backend/open_webui/tools/builtin.py b/backend/open_webui/tools/builtin.py
index 1823e71df2..61cd5ede4e 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,
@@ -760,14 +760,26 @@ async def search_notes(
content_snippet = ''
if note.data and note.data.get('content', {}).get('md'):
md_content = note.data['content']['md']
- lower_content = md_content.lower()
- lower_query = query.lower()
- idx = lower_content.find(lower_query)
- if idx != -1:
- start = max(0, idx - 50)
- end = min(len(md_content), idx + len(query) + 100)
+ content_lower = md_content.lower()
+
+ # Find the first matching word to center the snippet around.
+ search_words = query.lower().split()
+ match_pos = -1
+ match_len = len(query)
+ for word in search_words:
+ found_pos = content_lower.find(word)
+ if found_pos != -1:
+ match_pos = found_pos
+ match_len = len(word)
+ break
+
+ if match_pos != -1:
+ snippet_start = max(0, match_pos - 50)
+ snippet_end = min(len(md_content), match_pos + match_len + 100)
content_snippet = (
- ('...' if start > 0 else '') + md_content[start:end] + ('...' if end < len(md_content) else '')
+ ('...' if snippet_start > 0 else '')
+ + md_content[snippet_start:snippet_end]
+ + ('...' if snippet_end < len(md_content) else '')
)
else:
content_snippet = md_content[:150] + ('...' if len(md_content) > 150 else '')
@@ -808,18 +820,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 +890,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 +933,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 +959,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 +1010,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 +1085,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 +1157,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 +1213,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}
@@ -1213,7 +1225,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,18 +1286,18 @@ 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'})
# 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:
@@ -1336,24 +1348,24 @@ 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'})
# 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:
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 = []
@@ -1427,9 +1439,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 +1454,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 +1498,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 +1513,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 +1564,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 +1589,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 +1606,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 +1629,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 +1645,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 +1653,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 +1731,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 +1741,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 +1823,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 +1838,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 +1915,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 +1926,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 +1938,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 +1955,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 +1966,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 +2048,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 +2065,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 +2081,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 +2111,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 +2126,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 +2205,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 +2215,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 +2259,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 +2307,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 +2315,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,
@@ -2337,13 +2349,40 @@ VALID_TASK_STATUSES = {'pending', 'in_progress', 'completed', 'cancelled'}
class TaskItem(BaseModel):
id: Optional[str] = Field(None, description='Unique identifier for the task. Auto-generated if omitted.')
- content: Optional[str] = Field(None, description='Task description. Aliases: title, name, description.')
+ content: str = Field(..., description='Task description.')
status: Literal['pending', 'in_progress', 'completed', 'cancelled'] = Field('pending', description='Task status.')
-async def tasks(
- tasks: Optional[list[TaskItem]] = None,
- overwrite: bool = True,
+def _task_summary(all_tasks: list[dict]) -> dict:
+ """Build summary counts for a task list."""
+ pending = sum(1 for t in all_tasks if t['status'] == 'pending')
+ in_progress = sum(1 for t in all_tasks if t['status'] == 'in_progress')
+ completed = sum(1 for t in all_tasks if t['status'] == 'completed')
+ cancelled = sum(1 for t in all_tasks if t['status'] == 'cancelled')
+ return {
+ 'total': len(all_tasks),
+ 'pending': pending,
+ 'in_progress': in_progress,
+ 'completed': completed,
+ 'cancelled': cancelled,
+ }
+
+
+async def _emit_tasks(event_emitter, all_tasks: list[dict]):
+ """Persist task state to the UI."""
+ if event_emitter:
+ await event_emitter(
+ {
+ 'type': 'chat:message:tasks',
+ 'data': {
+ 'tasks': all_tasks,
+ },
+ }
+ )
+
+
+async def create_tasks(
+ tasks: list[TaskItem],
__chat_id__: str = None,
__message_id__: str = None,
__event_emitter__: callable = None,
@@ -2351,146 +2390,425 @@ async def tasks(
__user__: dict = None,
) -> str:
"""
- Track progress on multi-step work by maintaining a task checklist.
- Use this whenever a request involves multiple steps or could take
- significant effort. Call to set the full list, then call again
- with overwrite=false after completing each task to mark it
- completed. Do not leave tasks in_progress when the work is done.
- Each task has an id, content, and status (pending, in_progress,
- completed, cancelled).
+ Create a task checklist to track progress on multi-step work.
+ Call this once at the start to define all steps, then use
+ update_task to mark each task as you complete it.
- :param tasks: Optional list of task items. Each item: id (string), content (string, required for new tasks), status (pending|in_progress|completed|cancelled). Leave empty to fetch without modifying.
- :param overwrite: If true (default), replaces the entire task list. If false, updates/adds tasks by id while keeping existing ones.
+ :param tasks: List of task items. Each item: content (string, required), status (pending|in_progress|completed|cancelled, default pending), id (optional, auto-generated).
:return: JSON with the full task list and summary counts
"""
if __chat_id__ is None:
return json.dumps({'error': 'Chat context not available'})
try:
-
- def _to_dict(task) -> dict:
- """Convert TaskItem or dict to plain dict."""
+ all_tasks = []
+ for idx, task in enumerate(tasks):
if hasattr(task, 'model_dump'):
d = task.model_dump(exclude_none=True)
- # Include any extra fields the model sent
- if hasattr(task, 'model_extra') and task.model_extra:
- d.update(task.model_extra)
- return d
- return dict(task) if not isinstance(task, dict) else task
+ elif isinstance(task, dict):
+ d = task
+ else:
+ d = dict(task)
- def _resolve_content(d: dict) -> str:
- """Accept content, title, name, or description as the task text."""
- for key in ('content', 'title', 'name', 'description'):
- val = str(d.get(key, '')).strip()
- if val:
- return val
- return ''
+ content = str(d.get('content', '')).strip()
+ if not content:
+ continue
- def _resolve_id(d: dict, idx: int) -> str:
- """Use provided id, or auto-generate from index."""
- item_id = str(d.get('id', '') or '').strip()
- return item_id if item_id else str(idx + 1)
+ item_id = str(d.get('id', '') or '').strip() or str(idx + 1)
+ status = str(d.get('status', 'pending')).strip().lower()
+ if status not in VALID_TASK_STATUSES:
+ status = 'pending'
- if tasks is None:
- # Read-only - return current list
- all_tasks = Chats.get_chat_tasks_by_id(__chat_id__)
- elif overwrite:
- # Full replacement - validate and write
- all_tasks = []
- for idx, task in enumerate(tasks):
- d = _to_dict(task)
- item_id = _resolve_id(d, idx)
- content = _resolve_content(d)
- if not content:
- continue
+ all_tasks.append({'id': item_id, 'content': content, 'status': status})
- status = str(d.get('status', 'pending')).strip().lower()
- if status not in VALID_TASK_STATUSES:
- status = 'pending'
-
- all_tasks.append(
- {
- 'id': item_id,
- 'content': content,
- 'status': status,
- }
- )
- else:
- # Partial update - merge by id
- existing_tasks = Chats.get_chat_tasks_by_id(__chat_id__)
- existing_by_id = {t['id']: t for t in existing_tasks}
-
- seen_ids = set()
- for idx, task in enumerate(tasks):
- d = _to_dict(task)
- item_id = _resolve_id(d, len(existing_tasks) + idx)
-
- seen_ids.add(item_id)
-
- if item_id in existing_by_id:
- resolved = _resolve_content(d)
- if resolved:
- existing_by_id[item_id]['content'] = resolved
- status = str(d.get('status', '')).strip().lower()
- if status and status in VALID_TASK_STATUSES:
- existing_by_id[item_id]['status'] = status
- else:
- content = _resolve_content(d)
- if not content:
- continue
-
- status = str(d.get('status', 'pending')).strip().lower()
- if status not in VALID_TASK_STATUSES:
- status = 'pending'
-
- existing_by_id[item_id] = {
- 'id': item_id,
- 'content': content,
- 'status': status,
- }
-
- # Preserve order of existing, append new
- all_tasks = []
- for t in existing_tasks:
- if t['id'] in existing_by_id:
- all_tasks.append(existing_by_id[t['id']])
- for item_id in seen_ids:
- if not any(t['id'] == item_id for t in existing_tasks):
- all_tasks.append(existing_by_id[item_id])
-
- # Persist to DB and emit (skip for read-only)
- if tasks is not None:
- Chats.update_chat_tasks_by_id(__chat_id__, all_tasks)
-
- if __event_emitter__:
- await __event_emitter__(
- {
- 'type': 'chat:message:tasks',
- 'data': {
- 'tasks': all_tasks,
- },
- }
- )
-
- # Build summary counts
- pending = sum(1 for t in all_tasks if t['status'] == 'pending')
- in_progress = sum(1 for t in all_tasks if t['status'] == 'in_progress')
- completed = sum(1 for t in all_tasks if t['status'] == 'completed')
- cancelled = sum(1 for t in all_tasks if t['status'] == 'cancelled')
+ await Chats.update_chat_tasks_by_id(__chat_id__, all_tasks)
+ await _emit_tasks(__event_emitter__, all_tasks)
return json.dumps(
- {
- 'tasks': all_tasks,
- 'summary': {
- 'total': len(all_tasks),
- 'pending': pending,
- 'in_progress': in_progress,
- 'completed': completed,
- 'cancelled': cancelled,
- },
- },
+ {'tasks': all_tasks, 'summary': _task_summary(all_tasks)},
ensure_ascii=False,
)
except Exception as e:
log.exception(f'tasks error: {e}')
return json.dumps({'error': str(e)})
+
+
+async def update_task(
+ id: str,
+ status: str = 'completed',
+ __chat_id__: str = None,
+ __message_id__: str = None,
+ __event_emitter__: callable = None,
+ __request__: Request = None,
+ __user__: dict = None,
+) -> str:
+ """
+ Mark a single task as completed, in_progress, pending, or cancelled.
+ Call this after finishing each step. You MUST call this for every
+ task, including the very last one.
+
+ :param id: The task ID to update
+ :param status: New status: completed, in_progress, pending, or cancelled (default: completed)
+ :return: JSON with the updated task list and summary counts
+ """
+ if __chat_id__ is None:
+ return json.dumps({'error': 'Chat context not available'})
+
+ try:
+ status = status.strip().lower()
+ if status not in VALID_TASK_STATUSES:
+ return json.dumps(
+ {'error': f'Invalid status: {status}. Must be one of: {", ".join(sorted(VALID_TASK_STATUSES))}'}
+ )
+
+ all_tasks = await Chats.get_chat_tasks_by_id(__chat_id__)
+
+ found = False
+ for task in all_tasks:
+ if task['id'] == id:
+ task['status'] = status
+ found = True
+ break
+
+ if not found:
+ return json.dumps({'error': f'Task with id "{id}" not found'})
+
+ await Chats.update_chat_tasks_by_id(__chat_id__, all_tasks)
+ await _emit_tasks(__event_emitter__, all_tasks)
+
+ return json.dumps(
+ {'tasks': all_tasks, 'summary': _task_summary(all_tasks)},
+ ensure_ascii=False,
+ )
+ except Exception as e:
+ log.exception(f'update_task_status 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 = await 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 = await 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 = await Users.get_user_by_id(user_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'})
+
+ # 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 = await 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 = await Users.get_user_by_id(user_id)
+
+ result = await 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 = await Users.get_user_by_id(user_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 = await 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 = 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
+ await AutomationRuns.delete_by_automation(automation_id)
+ await 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/access_control/__init__.py b/backend/open_webui/utils/access_control/__init__.py
index f31c59e158..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,
@@ -255,3 +255,47 @@ def filter_allowed_access_grants(
access_grants = strip_user_access_grants(access_grants)
return access_grants
+
+
+async 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 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).
+ """
+ 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 await Groups.get_groups_by_member_id(user.id)}
+ if not (
+ user.id == model_info.user_id
+ or await 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')
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/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 = {
diff --git a/backend/open_webui/utils/auth.py b/backend/open_webui/utils/auth.py
index 16bc36500b..e0f331a9df 100644
--- a/backend/open_webui/utils/auth.py
+++ b/backend/open_webui/utils/auth.py
@@ -19,8 +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
from open_webui.models.users import Users
@@ -30,6 +28,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,
@@ -238,9 +237,7 @@ async def is_valid_token(request, decoded) -> bool:
# Per-user revocation (OIDC back-channel logout)
user_id = decoded.get('id')
if user_id:
- revoked_at = await request.app.state.redis.get(
- f'{REDIS_KEY_PREFIX}:auth:user:{user_id}:revoked_at'
- )
+ revoked_at = await request.app.state.redis.get(f'{REDIS_KEY_PREFIX}:auth:user:{user_id}:revoked_at')
if revoked_at:
try:
revoked_at_ts = int(revoked_at)
@@ -324,15 +321,18 @@ 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
- 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
@@ -353,7 +353,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,
@@ -369,17 +369,21 @@ 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
- # 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)
+ 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
+ # 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(
@@ -401,9 +405,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(
@@ -413,7 +417,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,
@@ -421,15 +425,33 @@ 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')
+ # 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,
+ )
- Users.update_last_active_by_id(user.id)
+ # Add user info to current span
+ 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')
+
+ await Users.update_last_active_by_id(user.id)
return user
@@ -451,7 +473,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.
@@ -461,14 +483,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 4d9eb2fb6c..5997b0e58b 100644
--- a/backend/open_webui/utils/automations.py
+++ b/backend/open_webui/utils/automations.py
@@ -26,7 +26,7 @@ from open_webui.models.automations import Automations, AutomationRuns, Automatio
from open_webui.models.chats import ChatForm, Chats
from open_webui.models.users import Users
from open_webui.utils.task import prompt_template
-from open_webui.internal.db import get_db
+from open_webui.internal.db import get_async_db
log = logging.getLogger(__name__)
@@ -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
############################
@@ -106,8 +125,8 @@ async def automation_worker_loop(app) -> None:
log.info(f'Automation worker started (poll interval: {AUTOMATION_POLL_INTERVAL}s)')
while True:
try:
- with get_db() as db:
- batch = Automations.claim_due(int(time.time_ns()), limit=10, db=db)
+ async with get_async_db() as 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:
@@ -205,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.
@@ -264,9 +293,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)
@@ -278,7 +307,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={
@@ -317,7 +346,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
@@ -338,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 = {
@@ -385,11 +409,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])
####################
@@ -397,12 +421,12 @@ async def execute_automation(app, automation: AutomationModel) -> None:
####################
-def _record_run(
+async def _record_run(
automation_id: str,
status: str,
chat_id: str = None,
error: str = None,
):
"""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)
+ async with get_async_db() as 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..21e3af5752 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''
- 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'')
+ 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/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/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..4d4a3726c3 100644
--- a/backend/open_webui/utils/middleware.py
+++ b/backend/open_webui/utils/middleware.py
@@ -889,10 +889,14 @@ def get_source_context(sources: list, source_ids: dict = None, include_content:
if source_id not in source_ids:
source_ids[source_id] = len(source_ids) + 1
src_name = source.get('source', {}).get('name')
+ src_type = source.get('source', {}).get('type')
+ src_rid = source.get('source', {}).get('id')
body = doc if include_content else ''
context_string += (
f'`` sources.
+_SAFE_STATIC_PATHS = frozenset(
+ {
+ '/user.png',
+ '/favicon.png',
+ '/static/favicon.png',
+ }
+)
def validate_profile_image_url(url: str) -> str:
@@ -16,28 +34,49 @@ def validate_profile_image_url(url: str) -> str:
Allowed formats:
- Empty string (falls back to default avatar)
- - data:image/* URIs (base64-encoded uploads from the frontend)
- - Known static asset paths (/user.png, /static/favicon.png)
- - Trusted external URLs (e.g. Gravatar)
+ - Known static-asset paths assigned by OWUI (exact match)
+ - The OWUI profile-image API route ``/api/v1/users/{id}/profile/image``
+ - ``http://`` and ``https://`` URLs with a valid hostname
+ - ``data:image/{png,jpeg,gif,webp};base64,...`` URIs
- Returns the url unchanged if valid, raises ValueError otherwise.
+ Everything else is rejected, including:
+ - Dangerous schemes (javascript:, file:, ftp:, …)
+ - SVG data URIs (can contain embedded scripts)
+ - Arbitrary relative paths (prevents authenticated GET triggers)
+ - Scheme-relative URLs (``//host/path``)
"""
if not url:
return url
- _ALLOWED_DATA_PREFIXES = (
- 'data:image/png',
- 'data:image/jpeg',
- 'data:image/gif',
- 'data:image/webp',
+ # --- Relative paths (exact match + anchored regex only) -----------
+
+ if url in _SAFE_STATIC_PATHS:
+ return url
+
+ if _USER_PROFILE_IMAGE_RE.match(url):
+ return url
+
+ # --- Absolute URLs -------------------------------------------------
+
+ # urlparse normalises the scheme to lowercase, giving us
+ # case-insensitive scheme matching for free.
+ parsed = urlparse(url)
+
+ # External images served over HTTP(S), e.g. OAuth provider avatars.
+ # Require a non-empty hostname (not just netloc, which can be ":80"
+ # for a URL like http://:80/path with no actual host).
+ if parsed.scheme in ('http', 'https'):
+ if not parsed.hostname:
+ raise ValueError('Invalid profile image URL: HTTP(S) URLs must include a host.')
+ return url
+
+ # Base64-encoded raster images uploaded via the frontend.
+ # The regex enforces the ;base64, boundary and is case-insensitive
+ # per the data-URI / MIME-type specs.
+ if _SAFE_DATA_URI_RE.match(url):
+ return url
+
+ raise ValueError(
+ 'Invalid profile image URL: must be a known internal path, '
+ 'an HTTP(S) URL with a host, or a data:image URI (png/jpeg/gif/webp).'
)
- if any(url.startswith(prefix) for prefix in _ALLOWED_DATA_PREFIXES):
- return url
-
- if url in _ALLOWED_STATIC_PATHS:
- return url
-
- if any(url.startswith(prefix) for prefix in _ALLOWED_URL_PREFIXES):
- return url
-
- raise ValueError('Invalid profile image URL: only data URIs and default avatars are allowed.')
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
diff --git a/src/lib/apis/index.ts b/src/lib/apis/index.ts
index 05475b4b37..5faf56d4d4 100644
--- a/src/lib/apis/index.ts
+++ b/src/lib/apis/index.ts
@@ -273,10 +273,42 @@ export const stopTask = async (token: string, id: string) => {
return res;
};
+export const stopTasksByChatId = async (token: string, chat_id: string) => {
+ let error = null;
+
+ const res = await fetch(`${WEBUI_BASE_URL}/api/tasks/chat/${encodeURIComponent(chat_id)}/stop`, {
+ method: 'POST',
+ headers: {
+ Accept: 'application/json',
+ 'Content-Type': 'application/json',
+ ...(token && { authorization: `Bearer ${token}` })
+ }
+ })
+ .then(async (res) => {
+ if (!res.ok) throw await res.json();
+ return res.json();
+ })
+ .catch((err) => {
+ console.error(err);
+ if ('detail' in err) {
+ error = err.detail;
+ } else {
+ error = err;
+ }
+ return null;
+ });
+
+ if (error) {
+ throw error;
+ }
+
+ return res;
+};
+
export const getTaskIdsByChatId = async (token: string, chat_id: string) => {
let error = null;
- const res = await fetch(`${WEBUI_BASE_URL}/api/tasks/chat/${chat_id}`, {
+ const res = await fetch(`${WEBUI_BASE_URL}/api/tasks/chat/${encodeURIComponent(chat_id)}`, {
method: 'GET',
headers: {
Accept: 'application/json',
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 @@
{formatJSONString(
+ args
+ )}
{JSON.stringify(
+ parsedResult,
+ null,
+ 2
+ )}
{:else}
{@const resultStr = String(parsedResult)}
{@const isTruncated = resultStr.length > RESULT_PREVIEW_LIMIT && !expandedResult}
diff --git a/src/lib/components/playground/Completions.svelte b/src/lib/components/playground/Completions.svelte
index a59f1d6747..e69aae35da 100644
--- a/src/lib/components/playground/Completions.svelte
+++ b/src/lib/components/playground/Completions.svelte
@@ -10,7 +10,6 @@
import { splitStream } from '$lib/utils';
import Spinner from '$lib/components/common/Spinner.svelte';
-
const i18n = getContext('i18n');
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')
}
};
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 @@