diff --git a/README.md b/README.md index a178b3271e..1885f4f6f1 100644 --- a/README.md +++ b/README.md @@ -172,8 +172,6 @@ After installation, you can access Open WebUI at [http://localhost:3000](http:// We offer various installation alternatives, including non-Docker native installation methods, Docker Compose, Kustomize, and Helm. Visit our [Open WebUI Documentation](https://docs.openwebui.com/getting-started/) or join our [Discord community](https://discord.gg/5rJgQTnV4s) for comprehensive guidance. -Look at the [Local Development Guide](https://docs.openwebui.com/getting-started/development) for instructions on setting up a local development environment. - ### Troubleshooting Encountering connection issues? Our [Open WebUI Documentation](https://docs.openwebui.com/troubleshooting/) has got you covered. For further assistance and to join our vibrant community, visit the [Open WebUI Discord](https://discord.gg/5rJgQTnV4s). diff --git a/backend/open_webui/config.py b/backend/open_webui/config.py index 58810c9e3e..fe25dda1c1 100644 --- a/backend/open_webui/config.py +++ b/backend/open_webui/config.py @@ -687,7 +687,6 @@ def load_oauth_providers(): return client OAUTH_PROVIDERS['google'] = { - 'redirect_uri': GOOGLE_REDIRECT_URI.value, 'register': google_oauth_register, } @@ -708,7 +707,6 @@ def load_oauth_providers(): return client OAUTH_PROVIDERS['microsoft'] = { - 'redirect_uri': MICROSOFT_REDIRECT_URI.value, 'picture_url': MICROSOFT_CLIENT_PICTURE_URL.value, 'register': microsoft_oauth_register, } @@ -733,7 +731,6 @@ def load_oauth_providers(): return client OAUTH_PROVIDERS['github'] = { - 'redirect_uri': GITHUB_CLIENT_REDIRECT_URI.value, 'register': github_oauth_register, 'sub_claim': 'id', } @@ -775,7 +772,6 @@ def load_oauth_providers(): OAUTH_PROVIDERS['oidc'] = { 'name': OAUTH_PROVIDER_NAME.value, - 'redirect_uri': OPENID_REDIRECT_URI.value, 'register': oidc_oauth_register, } @@ -919,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) @@ -1226,10 +1223,16 @@ DEFAULT_MODEL_METADATA = PersistentConfig( {}, ) +try: + default_model_params = json.loads(os.environ.get('DEFAULT_MODEL_PARAMS', '{}')) +except Exception as e: + log.exception(f'Error loading DEFAULT_MODEL_PARAMS: {e}') + default_model_params = {} + DEFAULT_MODEL_PARAMS = PersistentConfig( 'DEFAULT_MODEL_PARAMS', 'models.default_params', - {}, + default_model_params, ) DEFAULT_USER_ROLE = PersistentConfig( @@ -1435,6 +1438,10 @@ USER_PERMISSIONS_FEATURES_API_KEYS = os.environ.get('USER_PERMISSIONS_FEATURES_A USER_PERMISSIONS_FEATURES_MEMORIES = os.environ.get('USER_PERMISSIONS_FEATURES_MEMORIES', 'True').lower() == 'true' +USER_PERMISSIONS_FEATURES_AUTOMATIONS = ( + os.environ.get('USER_PERMISSIONS_FEATURES_AUTOMATIONS', 'False').lower() == 'true' +) + USER_PERMISSIONS_SETTINGS_INTERFACE = os.environ.get('USER_PERMISSIONS_SETTINGS_INTERFACE', 'True').lower() == 'true' @@ -1504,6 +1511,7 @@ DEFAULT_USER_PERMISSIONS = { 'image_generation': USER_PERMISSIONS_FEATURES_IMAGE_GENERATION, 'code_interpreter': USER_PERMISSIONS_FEATURES_CODE_INTERPRETER, 'memories': USER_PERMISSIONS_FEATURES_MEMORIES, + 'automations': USER_PERMISSIONS_FEATURES_AUTOMATIONS, }, 'settings': { 'interface': USER_PERMISSIONS_SETTINGS_INTERFACE, @@ -1534,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', @@ -3084,7 +3104,7 @@ WEB_SEARCH_CONCURRENT_REQUESTS = PersistentConfig( WEB_FETCH_MAX_CONTENT_LENGTH = PersistentConfig( 'WEB_FETCH_MAX_CONTENT_LENGTH', - 'rag.web.search.fetch_url_max_content_length', + 'rag.web.fetch.max_content_length', (int(os.environ.get('WEB_FETCH_MAX_CONTENT_LENGTH')) if os.environ.get('WEB_FETCH_MAX_CONTENT_LENGTH') else None), ) @@ -3933,6 +3953,18 @@ AUDIO_TTS_AZURE_SPEECH_OUTPUT_FORMAT = PersistentConfig( os.getenv('AUDIO_TTS_AZURE_SPEECH_OUTPUT_FORMAT', 'audio-24khz-160kbitrate-mono-mp3'), ) +AUDIO_TTS_MISTRAL_API_KEY = PersistentConfig( + 'AUDIO_TTS_MISTRAL_API_KEY', + 'audio.tts.mistral.api_key', + os.getenv('AUDIO_TTS_MISTRAL_API_KEY', ''), +) + +AUDIO_TTS_MISTRAL_API_BASE_URL = PersistentConfig( + 'AUDIO_TTS_MISTRAL_API_BASE_URL', + 'audio.tts.mistral.api_base_url', + os.getenv('AUDIO_TTS_MISTRAL_API_BASE_URL', 'https://api.mistral.ai/v1'), +) + #################################### # LDAP diff --git a/backend/open_webui/env.py b/backend/open_webui/env.py index f0dbf9c114..5e1e150d04 100644 --- a/backend/open_webui/env.py +++ b/backend/open_webui/env.py @@ -60,13 +60,14 @@ if USE_CUDA.lower() == 'true': else: DEVICE_TYPE = 'cpu' -try: - import torch +if sys.platform == 'darwin': + try: + import torch - if torch.backends.mps.is_available() and torch.backends.mps.is_built(): - DEVICE_TYPE = 'mps' -except Exception: - pass + if torch.backends.mps.is_available() and torch.backends.mps.is_built(): + DEVICE_TYPE = 'mps' + except Exception: + pass #################################### # LOGGING @@ -425,6 +426,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 == '': @@ -495,6 +517,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) #################################### @@ -544,6 +571,12 @@ OAUTH_MAX_SESSIONS_PER_USER = int(os.environ.get('OAUTH_MAX_SESSIONS_PER_USER', # Allows external apps to exchange OAuth tokens for OpenWebUI tokens ENABLE_OAUTH_TOKEN_EXCHANGE = os.environ.get('ENABLE_OAUTH_TOKEN_EXCHANGE', 'False').lower() == 'true' +# Back-Channel Logout Configuration +# When enabled, exposes POST /oauth/backchannel-logout for IdP-initiated logout +# per OpenID Connect Back-Channel Logout 1.0 spec. +# Requires Redis for JWT revocation. +ENABLE_OAUTH_BACKCHANNEL_LOGOUT = os.environ.get('ENABLE_OAUTH_BACKCHANNEL_LOGOUT', 'False').lower() == 'true' + #################################### # SCIM Configuration #################################### @@ -771,6 +804,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 == '': @@ -874,6 +937,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 9bfe77c41e..37a9011aab 100644 --- a/backend/open_webui/functions.py +++ b/backend/open_webui/functions.py @@ -34,7 +34,6 @@ from open_webui.utils.plugin import ( load_function_module_by_id, get_function_module_from_cache, ) -from open_webui.utils.tools import get_tools from open_webui.env import GLOBAL_LOG_LEVEL @@ -54,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: @@ -74,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'): @@ -188,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 = {} @@ -199,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: @@ -209,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', {}) @@ -226,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) @@ -255,17 +254,7 @@ async def generate_function_chat_completion(request, form_data, user, models: di '__oauth_token__': oauth_token, '__request__': request, } - extra_params['__tools__'] = await get_tools( - request, - tool_ids, - user, - { - **extra_params, - '__model__': models.get(form_data['model'], None), - '__messages__': form_data['messages'], - '__files__': files, - }, - ) + extra_params['__tools__'] = metadata.get('tools', {}) if model_info: if model_info.base_model_id: @@ -279,10 +268,10 @@ async def generate_function_chat_completion(request, form_data, user, models: di form_data = apply_system_prompt_to_body(system, form_data, metadata, user) pipe_id = get_pipe_id(form_data) - function_module = get_function_module_by_id(request, pipe_id) + function_module = await get_function_module_by_id(request, pipe_id) pipe = function_module.pipe - params = get_function_params(function_module, form_data, user, extra_params) + params = await get_function_params(function_module, form_data, user, extra_params) if form_data.get('stream', False): diff --git a/backend/open_webui/internal/db.py b/backend/open_webui/internal/db.py index b0545255a6..3818543fc7 100644 --- a/backend/open_webui/internal/db.py +++ b/backend/open_webui/internal/db.py @@ -1,7 +1,7 @@ import os import json import logging -from contextlib import contextmanager +from contextlib import asynccontextmanager, contextmanager from typing import Any, Optional from open_webui.internal.wrappers import register_connection @@ -19,6 +19,7 @@ from open_webui.env import ( ) from peewee_migrate import Router from sqlalchemy import Dialect, create_engine, MetaData, event, types +from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker from sqlalchemy.ext.declarative import declarative_base from sqlalchemy.orm import scoped_session, sessionmaker, Session from sqlalchemy.pool import QueuePool, NullPool @@ -81,6 +82,32 @@ if ENABLE_DB_MIGRATIONS: SQLALCHEMY_DATABASE_URL = DATABASE_URL + +def _make_async_url(url: str) -> str: + """Convert a sync database URL to its async driver equivalent.""" + if url.startswith('sqlite+sqlcipher://'): + # SQLCipher has no async driver — not supported for async + raise ValueError( + 'sqlite+sqlcipher:// URLs are not supported with async engine. ' + 'Use standard sqlite:// or postgresql:// instead.' + ) + if url.startswith('sqlite:///') or url.startswith('sqlite://'): + return url.replace('sqlite://', 'sqlite+aiosqlite://', 1) + if url.startswith('postgresql+psycopg2://'): + return url.replace('postgresql+psycopg2://', 'postgresql+asyncpg://', 1) + if url.startswith('postgresql://'): + return url.replace('postgresql://', 'postgresql+asyncpg://', 1) + if url.startswith('postgres://'): + return url.replace('postgres://', 'postgresql+asyncpg://', 1) + # For other dialects, return as-is and let SQLAlchemy handle it + return url + + +# ============================================================ +# SYNC ENGINE (used only for: startup migrations, config loading, +# Alembic, peewee migration, health checks) +# ============================================================ + # Handle SQLCipher URLs if SQLALCHEMY_DATABASE_URL.startswith('sqlite+sqlcipher://'): database_password = os.environ.get('DATABASE_PASSWORD') @@ -155,6 +182,7 @@ else: engine = create_engine(SQLALCHEMY_DATABASE_URL, pool_pre_ping=True) +# Sync session — used ONLY for startup config loading (config.py runs at import time) SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine, expire_on_commit=False) metadata_obj = MetaData(schema=DATABASE_SCHEMA) Base = declarative_base(metadata=metadata_obj) @@ -162,6 +190,7 @@ ScopedSession = scoped_session(SessionLocal) def get_session(): + """Sync session generator — used ONLY for startup/config operations.""" db = SessionLocal() try: yield db @@ -172,10 +201,82 @@ def get_session(): get_db = contextmanager(get_session) -@contextmanager -def get_db_context(db: Optional[Session] = None): - if isinstance(db, Session) and DATABASE_ENABLE_SESSION_SHARING: +# ============================================================ +# ASYNC ENGINE (used for ALL runtime database operations) +# ============================================================ + +ASYNC_SQLALCHEMY_DATABASE_URL = _make_async_url(SQLALCHEMY_DATABASE_URL) + +if 'sqlite' in ASYNC_SQLALCHEMY_DATABASE_URL: + async_engine = create_async_engine( + ASYNC_SQLALCHEMY_DATABASE_URL, + connect_args={'check_same_thread': False}, + ) + + if DATABASE_ENABLE_SQLITE_WAL: + + @event.listens_for(async_engine.sync_engine, 'connect') + def _set_sqlite_wal(dbapi_connection, connection_record): + cursor = dbapi_connection.cursor() + cursor.execute('PRAGMA journal_mode=WAL') + cursor.close() +else: + if isinstance(DATABASE_POOL_SIZE, int): + if DATABASE_POOL_SIZE > 0: + async_engine = create_async_engine( + ASYNC_SQLALCHEMY_DATABASE_URL, + pool_size=DATABASE_POOL_SIZE, + max_overflow=DATABASE_POOL_MAX_OVERFLOW, + pool_timeout=DATABASE_POOL_TIMEOUT, + pool_recycle=DATABASE_POOL_RECYCLE, + pool_pre_ping=True, + ) + 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 4f1b6551d7..975b38a94b 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, @@ -97,6 +98,7 @@ from open_webui.routers import ( utils, scim, terminals, + automations, ) from open_webui.routers.retrieval import ( @@ -107,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 @@ -211,6 +213,8 @@ from open_webui.config import ( AUDIO_TTS_AZURE_SPEECH_REGION, AUDIO_TTS_AZURE_SPEECH_BASE_URL, AUDIO_TTS_AZURE_SPEECH_OUTPUT_FORMAT, + AUDIO_TTS_MISTRAL_API_KEY, + AUDIO_TTS_MISTRAL_API_BASE_URL, PLAYWRIGHT_WS_URL, PLAYWRIGHT_TIMEOUT, FIRECRAWL_API_BASE_URL, @@ -380,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, @@ -470,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, @@ -510,6 +517,8 @@ from open_webui.env import ( WEBUI_ADMIN_NAME, ENABLE_EASTER_EGGS, LOG_FORMAT, + # OAuth Back-Channel Logout + ENABLE_OAUTH_BACKCHANNEL_LOGOUT, ) @@ -558,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 @@ -568,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__) @@ -622,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, @@ -648,6 +661,10 @@ async def lifespan(app: FastAPI): asyncio.create_task(periodic_usage_pool_cleanup()) asyncio.create_task(periodic_session_pool_cleanup()) + from open_webui.utils.automations import automation_worker_loop + + asyncio.create_task(automation_worker_loop(app)) + if app.state.config.ENABLE_BASE_MODELS_CACHE: try: await get_all_models( @@ -704,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() @@ -865,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 @@ -1277,6 +1301,9 @@ app.state.config.TTS_AZURE_SPEECH_REGION = AUDIO_TTS_AZURE_SPEECH_REGION app.state.config.TTS_AZURE_SPEECH_BASE_URL = AUDIO_TTS_AZURE_SPEECH_BASE_URL app.state.config.TTS_AZURE_SPEECH_OUTPUT_FORMAT = AUDIO_TTS_AZURE_SPEECH_OUTPUT_FORMAT +app.state.config.TTS_MISTRAL_API_KEY = AUDIO_TTS_MISTRAL_API_KEY +app.state.config.TTS_MISTRAL_API_BASE_URL = AUDIO_TTS_MISTRAL_API_BASE_URL + app.state.faster_whisper_model = None app.state.speech_synthesiser = None @@ -1371,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) @@ -1522,6 +1504,7 @@ if ENABLE_ADMIN_ANALYTICS: app.include_router(analytics.router, prefix='/api/v1/analytics', tags=['analytics']) app.include_router(utils.router, prefix='/api/v1/utils', tags=['utils']) app.include_router(terminals.router, prefix='/api/v1/terminals', tags=['terminals']) +app.include_router(automations.router, prefix='/api/v1/automations', tags=['automations']) # SCIM 2.0 API for identity management if ENABLE_SCIM: @@ -1540,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, ) ################################## @@ -1588,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])}' @@ -1654,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: @@ -1742,7 +1726,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, @@ -1754,7 +1738,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'), [ @@ -1786,7 +1770,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'], { @@ -1797,13 +1781,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'}, @@ -1814,12 +1798,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'], { @@ -1828,7 +1812,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', @@ -1842,17 +1826,32 @@ async def chat_completion( except Exception: pass finally: + # Clean up MCP clients. Shield the entire block from + # CancelledError so disconnect() can finish even when the + # task is being stopped. Each client is isolated so one + # failure doesn't skip the rest. try: if mcp_clients := metadata.get('mcp_clients'): - for client in reversed(mcp_clients.values()): - await client.disconnect() + + async def _cleanup_mcp(): + for client in reversed(list(mcp_clients.values())): + try: + await client.disconnect() + except Exception as e: + log.debug(f'Error disconnecting MCP client: {e}') + + await asyncio.wait_for( + asyncio.shield(_cleanup_mcp()), + timeout=10.0, + ) + except asyncio.TimeoutError: + log.warning('MCP client cleanup timed out after 10 s') except Exception as e: - log.debug(f'Error cleaning up: {e}') - pass + log.debug(f'Error cleaning up MCP clients: {e}') # 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: @@ -1866,7 +1865,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} @@ -1878,6 +1877,10 @@ async def chat_completion( generate_chat_completions = chat_completion generate_chat_completion = chat_completion +# Expose as app.state so internal callers (e.g. automations) can +# use the full pipeline without importing from main.py (avoids circular deps). +app.state.CHAT_COMPLETION_HANDLER = chat_completion + ################################## # @@ -1974,7 +1977,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 @@ -1983,15 +1986,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) @@ -1999,6 +2008,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 @@ -2030,9 +2054,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: @@ -2241,7 +2265,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 @@ -2448,11 +2472,26 @@ 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) +############################ +# OIDC Back-Channel Logout +############################ + + +@app.post('/oauth/backchannel-logout') +async def oauth_backchannel_logout( + request: Request, + db: AsyncSession = Depends(get_async_session), +): + if not ENABLE_OAUTH_BACKCHANNEL_LOGOUT: + raise HTTPException(status_code=404) + return await oauth_manager.handle_backchannel_logout(request, db=db) + + @app.get('/manifest.json') async def get_manifest_json(): if app.state.EXTERNAL_PWA_MANIFEST_URL: diff --git a/backend/open_webui/migrations/versions/b7c8d9e0f1a2_add_last_read_at_to_chat.py b/backend/open_webui/migrations/versions/b7c8d9e0f1a2_add_last_read_at_to_chat.py new file mode 100644 index 0000000000..fb254432f6 --- /dev/null +++ b/backend/open_webui/migrations/versions/b7c8d9e0f1a2_add_last_read_at_to_chat.py @@ -0,0 +1,27 @@ +"""add last_read_at to chat + +Revision ID: b7c8d9e0f1a2 +Revises: d4e5f6a7b8c9 +Create Date: 2026-04-01 04:00:00.000000 + +""" + +from alembic import op +import sqlalchemy as sa + + +# revision identifiers, used by Alembic. +revision = 'b7c8d9e0f1a2' +down_revision = 'd4e5f6a7b8c9' +branch_labels = None +depends_on = None + + +def upgrade(): + op.add_column('chat', sa.Column('last_read_at', sa.BigInteger(), nullable=True)) + # Set existing chats to be marked as read + op.execute('UPDATE chat SET last_read_at = updated_at') + + +def downgrade(): + op.drop_column('chat', 'last_read_at') diff --git a/backend/open_webui/migrations/versions/d4e5f6a7b8c9_add_automation_tables.py b/backend/open_webui/migrations/versions/d4e5f6a7b8c9_add_automation_tables.py new file mode 100644 index 0000000000..fc90dc417f --- /dev/null +++ b/backend/open_webui/migrations/versions/d4e5f6a7b8c9_add_automation_tables.py @@ -0,0 +1,55 @@ +"""add automation tables + +Revision ID: d4e5f6a7b8c9 +Revises: f1e2d3c4b5a6 +Create Date: 2026-03-30 +""" + +from typing import Union + +from alembic import op +import sqlalchemy as sa + +revision: str = 'd4e5f6a7b8c9' +down_revision: Union[str, None] = 'a3dd5bedd151' +branch_labels = None +depends_on = None + + +def upgrade(): + op.create_table( + 'automation', + sa.Column('id', sa.Text(), primary_key=True), + sa.Column('user_id', sa.Text(), nullable=False), + sa.Column('name', sa.Text(), nullable=False), + sa.Column('data', sa.JSON(), nullable=False), + sa.Column('meta', sa.JSON(), nullable=True), + sa.Column('is_active', sa.Boolean(), nullable=False, default=True), + sa.Column('last_run_at', sa.BigInteger(), nullable=True), + sa.Column('next_run_at', sa.BigInteger(), nullable=True), + sa.Column('created_at', sa.BigInteger(), nullable=False), + sa.Column('updated_at', sa.BigInteger(), nullable=False), + ) + op.create_index('ix_automation_next_run', 'automation', ['next_run_at']) + + op.create_table( + 'automation_run', + sa.Column('id', sa.Text(), primary_key=True), + sa.Column('automation_id', sa.Text(), nullable=False), + sa.Column('chat_id', sa.Text(), nullable=True), + sa.Column('status', sa.Text(), nullable=False), + sa.Column('error', sa.Text(), nullable=True), + sa.Column('created_at', sa.BigInteger(), nullable=False), + ) + op.create_index( + 'ix_automation_run_automation_id', + 'automation_run', + ['automation_id'], + ) + + +def downgrade(): + op.drop_index('ix_automation_run_automation_id') + op.drop_table('automation_run') + op.drop_index('ix_automation_next_run') + op.drop_table('automation') 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 new file mode 100644 index 0000000000..c891c3204e --- /dev/null +++ b/backend/open_webui/models/automations.py @@ -0,0 +1,377 @@ +import time +import logging +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, delete, update +from sqlalchemy.ext.asyncio import AsyncSession + +from open_webui.internal.db import Base, get_async_db_context + +log = logging.getLogger(__name__) + + +#################### +# Automation DB Schema +#################### + + +class Automation(Base): + __tablename__ = 'automation' + + id = Column(Text, primary_key=True) + user_id = Column(Text, nullable=False) + name = Column(Text, nullable=False) + data = Column(JSON, nullable=False) # {prompt, model_id, rrule} + meta = Column(JSON, nullable=True) + is_active = Column(Boolean, nullable=False, default=True) + last_run_at = Column(BigInteger, nullable=True) + next_run_at = Column(BigInteger, nullable=True) + + created_at = Column(BigInteger, nullable=False) + updated_at = Column(BigInteger, nullable=False) + + __table_args__ = (Index('ix_automation_next_run', 'next_run_at'),) + + +class AutomationRun(Base): + __tablename__ = 'automation_run' + + id = Column(Text, primary_key=True) + automation_id = Column(Text, nullable=False) + chat_id = Column(Text, nullable=True) + status = Column(Text, nullable=False) # success | error + error = Column(Text, nullable=True) + created_at = Column(BigInteger, nullable=False) + + __table_args__ = ( + Index('ix_automation_run_automation_id', 'automation_id'), + Index('ix_automation_run_aid_created', 'automation_id', 'created_at'), + ) + + +#################### +# Pydantic Models +#################### + + +class AutomationTerminalConfig(BaseModel): + server_id: str + cwd: Optional[str] = None + + +class AutomationData(BaseModel): + prompt: str + model_id: str + rrule: str + terminal: Optional[AutomationTerminalConfig] = None + + +class AutomationModel(BaseModel): + model_config = ConfigDict(from_attributes=True) + + id: str + user_id: str + name: str + data: dict + meta: Optional[dict] = None + is_active: bool + last_run_at: Optional[int] = None + next_run_at: Optional[int] = None + + created_at: int + updated_at: int + + +class AutomationRunModel(BaseModel): + model_config = ConfigDict(from_attributes=True) + + id: str + automation_id: str + chat_id: Optional[str] = None + status: str + error: Optional[str] = None + created_at: int + + +class AutomationForm(BaseModel): + name: str + data: AutomationData + meta: Optional[dict] = None + is_active: Optional[bool] = True + + +class AutomationResponse(AutomationModel): + last_run: Optional[AutomationRunModel] = None + next_runs: Optional[list[int]] = None + + +class AutomationListResponse(BaseModel): + items: list[AutomationModel] + total: int + + +#################### +# AutomationTable +#################### + + +class AutomationTable: + async def insert( + self, + user_id: str, + form: AutomationForm, + next_run_at: int, + db: Optional[AsyncSession] = None, + ) -> AutomationModel: + async with get_async_db_context(db) as db: + now = int(time.time_ns()) + row = Automation( + id=str(uuid4()), + user_id=user_id, + name=form.name, + data=form.data.model_dump(), + meta=form.meta, + is_active=form.is_active, + next_run_at=next_run_at, + created_at=now, + updated_at=now, + ) + db.add(row) + await db.commit() + await db.refresh(row) + return AutomationModel.model_validate(row) + + 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 + + async def search_automations( + self, + user_id: str, + query: Optional[str] = None, + status: Optional[str] = None, + skip: int = 0, + limit: int = 30, + db: Optional[AsyncSession] = None, + ) -> 'AutomationListResponse': + 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 + stmt = stmt.filter( + or_( + Automation.name.ilike(search), + cast(Automation.data, String).ilike(search), + ) + ) + + if status == 'active': + stmt = stmt.filter(Automation.is_active == True) + elif status == 'paused': + stmt = stmt.filter(Automation.is_active == False) + + stmt = stmt.order_by(Automation.created_at.desc()) + + # Get total count + count_result = await db.execute(select(func.count()).select_from(stmt.subquery())) + total = count_result.scalar() + + if skip: + stmt = stmt.offset(skip) + if limit: + stmt = stmt.limit(limit) + + result = await db.execute(stmt) + rows = result.scalars().all() + return AutomationListResponse( + items=[AutomationModel.model_validate(r) for r in rows], + total=total, + ) + + async def update_by_id( + self, + id: str, + form: AutomationForm, + next_run_at: int, + db: Optional[AsyncSession] = None, + ) -> Optional[AutomationModel]: + async with get_async_db_context(db) as db: + row = await db.get(Automation, id) + if not row: + return None + row.name = form.name + row.data = form.data.model_dump() + row.meta = form.meta + if form.is_active is not None: + row.is_active = form.is_active + row.next_run_at = next_run_at + row.updated_at = int(time.time_ns()) + await db.commit() + await db.refresh(row) + return AutomationModel.model_validate(row) + + async def toggle( + self, + id: str, + next_run_at: Optional[int], + db: Optional[AsyncSession] = None, + ) -> Optional[AutomationModel]: + async with get_async_db_context(db) as db: + row = await db.get(Automation, id) + if not row: + return None + row.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()) + await db.commit() + await db.refresh(row) + return AutomationModel.model_validate(row) + + 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 + await db.delete(row) + await db.commit() + return True + + async def claim_due(self, now_ns: int, limit: int = 10, db: Optional[AsyncSession] = None) -> list[AutomationModel]: + """ + Atomically claim due automations for execution. + + Advances next_run_at immediately so the row can never be + double-claimed. On PostgreSQL, uses FOR UPDATE SKIP LOCKED + for zero-contention distributed work claiming. + """ + async with get_async_db_context(db) as db: + stmt = ( + select(Automation) + .where( + Automation.is_active == True, + Automation.next_run_at <= now_ns, + ) + .order_by(Automation.next_run_at) + .limit(limit) + ) + + if db.bind.dialect.name == 'postgresql': + stmt = stmt.with_for_update(skip_locked=True) + + result = await db.execute(stmt) + rows = result.scalars().all() + + from open_webui.utils.automations import next_run_ns + + for row in rows: + row.last_run_at = now_ns + row.next_run_at = next_run_ns(row.data.get('rrule', '')) + + await db.commit() + + return [AutomationModel.model_validate(r) for r in rows] + + +#################### +# AutomationRunTable +#################### + + +class AutomationRunTable: + async def insert( + self, + automation_id: str, + status: str, + chat_id: Optional[str] = None, + error: Optional[str] = None, + db: Optional[AsyncSession] = None, + ) -> AutomationRunModel: + async with get_async_db_context(db) as db: + row = AutomationRun( + id=str(uuid4()), + automation_id=automation_id, + chat_id=chat_id, + status=status, + error=error, + created_at=int(time.time_ns()), + ) + db.add(row) + await db.commit() + await db.refresh(row) + return AutomationRunModel.model_validate(row) + + 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()) + .limit(1) + ) + row = result.scalars().first() + return AutomationRunModel.model_validate(row) if row else None + + async def get_latest_batch( + self, automation_ids: list[str], db: Optional[AsyncSession] = None + ) -> dict[str, AutomationRunModel]: + """Fetch the latest run for each automation in a single query.""" + if not automation_ids: + return {} + 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[AsyncSession] = None, + ) -> list[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()) + .offset(skip) + .limit(limit) + ) + rows = result.scalars().all() + return [AutomationRunModel.model_validate(r) for r in rows] + + 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() +AutomationRuns = AutomationRunTable() 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 97490c1602..bd9c720fa4 100644 --- a/backend/open_webui/models/chat_messages.py +++ b/backend/open_webui/models/chat_messages.py @@ -3,8 +3,10 @@ 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 from sqlalchemy import ( @@ -15,7 +17,6 @@ from sqlalchemy import ( Text, JSON, Index, - func, ) #################### @@ -41,6 +42,12 @@ def _normalize_timestamp(timestamp: int) -> float: return timestamp +def get_usage(data: dict) -> Optional[dict]: + """Extract and normalize usage from message data.""" + usage = data.get('usage') or (data.get('info') or {}).get('usage') + return normalize_usage(usage) if usage else None + + #################### # ChatMessage DB Schema #################### @@ -122,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: @@ -163,24 +170,21 @@ class ChatMessageTable: existing.status_history = data.get('status_history') or data.get('statusHistory') if 'error' in data: existing.error = data.get('error') - # Extract usage - check direct field first, then info.usage - usage = data.get('usage') - if not usage: - info = data.get('info', {}) - usage = info.get('usage') if info else None + # Extract and normalize usage + usage = get_usage(data) if usage: - existing.usage = usage + # Deep-merge: preserve existing keys not present in new data + # This prevents background tasks (follow-ups, title, tags) + # 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 - # Extract usage - check direct field first, then info.usage - usage = data.get('usage') - if not usage: - info = data.get('info', {}) - usage = info.get('usage') if info else None + # Extract and normalize usage + usage = get_usage(data) message = ChatMessage( id=composite_id, chat_id=chat_id, @@ -201,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, @@ -349,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'), @@ -362,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: { @@ -378,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, @@ -411,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'), @@ -424,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: { @@ -440,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]] = {} @@ -543,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 66d99b05cb..b13edbf0bf 100644 --- a/backend/open_webui/models/chats.py +++ b/backend/open_webui/models/chats.py @@ -4,11 +4,15 @@ 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 +from open_webui.models.automations import AutomationRun from open_webui.utils.misc import sanitize_data_for_db, sanitize_text_for_db from pydantic import BaseModel, ConfigDict @@ -23,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 @@ -57,17 +58,14 @@ class Chat(Base): tasks = Column(JSON, nullable=True) summary = Column(Text, nullable=True) + last_read_at = Column(BigInteger, nullable=True) + __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'), ) @@ -93,6 +91,8 @@ class ChatModel(BaseModel): tasks: Optional[list] = None summary: Optional[str] = None + last_read_at: Optional[int] = None + class ChatFile(Base): __tablename__ = 'chat_file' @@ -176,6 +176,7 @@ class ChatTitleIdResponse(BaseModel): title: str updated_at: int created_at: int + last_read_at: Optional[int] = None class SharedChatResponse(BaseModel): @@ -291,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( **{ @@ -310,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: @@ -319,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, @@ -347,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: @@ -361,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: @@ -370,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, @@ -381,35 +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_title_by_id(self, id: str, title: str) -> Optional[ChatModel]: - chat = self.get_chat_by_id(id) - if chat is None: + async def update_chat_last_read_at_by_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> bool: + try: + 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()) + await db.commit() + return True + return False + except Exception: + return False + + async def update_chat_title_by_id(self, id: str, title: str) -> Optional[ChatModel]: + try: + 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()) + await db.commit() + await db.refresh(chat_item) + return ChatModel.model_validate(chat_item) + except Exception: return None - chat = chat.chat - chat['title'] = title - - return self.update_chat_by_id(id, chat) - - 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 @@ -419,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 @@ -482,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, @@ -491,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 @@ -509,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 @@ -528,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( **{ @@ -557,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 @@ -580,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') @@ -682,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( { @@ -710,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') @@ -734,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( { @@ -771,80 +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, - ) -> list[ChatModel]: - with get_db_context(db) as db: - query = db.query(Chat).filter_by(user_id=user_id) + db: Optional[AsyncSession] = None, + ) -> list[ChatTitleIdResponse]: + 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) + 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() - return [ChatModel.model_validate(chat) for chat in all_chats] - - def get_chat_title_id_list_by_user_id( - self, - user_id: str, - include_archived: bool = False, - include_folders: bool = False, - include_pinned: bool = False, - skip: Optional[int] = None, - limit: Optional[int] = None, - db: Optional[Session] = 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 - ) - - if skip: - query = query.offset(skip) - if limit: - query = query.limit(limit) - - all_chats = query.all() - - # result has to be destructured from sqlalchemy `row` and mapped to a dict since the `ChatModel`is not the returned dataclass. + result = await db.execute(stmt) + all_chats = result.all() return [ ChatTitleIdResponse.model_validate( { @@ -852,111 +843,157 @@ class ChatTable: 'title': chat[1], 'updated_at': chat[2], 'created_at': chat[3], + 'last_read_at': chat[4], } ) for chat in all_chats ] - def get_chat_list_by_chat_ids( + async def get_chat_title_id_list_by_user_id( + self, + user_id: str, + include_archived: bool = False, + include_folders: bool = False, + include_pinned: bool = False, + skip: Optional[int] = None, + limit: Optional[int] = None, + db: Optional[AsyncSession] = None, + ) -> list[ChatTitleIdResponse]: + 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: + stmt = stmt.offset(skip) + if limit: + stmt = stmt.limit(limit) + + result = await db.execute(stmt) + all_chats = result.all() + + return [ + ChatTitleIdResponse.model_validate( + { + 'id': chat[0], + 'title': chat[1], + 'updated_at': chat[2], + 'created_at': chat[3], + 'last_read_at': chat[4], + } + ) + for chat in all_chats + ] + + 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') @@ -964,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( **{ @@ -989,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) ) + all_chats = result.all() return [ ChatTitleIdResponse.model_validate( { @@ -1004,24 +1045,27 @@ class ChatTable: 'title': chat[1], 'updated_at': chat[2], 'created_at': chat[3], + 'last_read_at': chat[4], } ) 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. @@ -1029,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:')], ) @@ -1077,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 = ( @@ -1111,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 @@ -1128,7 +1175,7 @@ class ChatTable: """) ) elif tag_ids: - query = query.filter( + stmt = stmt.filter( and_( *[ text(f""" @@ -1144,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 ( @@ -1164,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 @@ -1182,7 +1225,7 @@ class ChatTable: """) ) elif tag_ids: - query = query.filter( + stmt = stmt.filter( and_( *[ text(f""" @@ -1197,146 +1240,192 @@ 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, - ) -> list[ChatModel]: - 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) + db: Optional[AsyncSession] = None, + ) -> list[ChatTitleIdResponse]: + 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() - return [ChatModel.model_validate(chat) for chat in all_chats] + result = await db.execute(stmt) + all_chats = result.all() + return [ + ChatTitleIdResponse.model_validate( + { + 'id': chat[0], + 'title': chat[1], + 'updated_at': chat[2], + 'created_at': chat[3], + 'last_read_at': chat[4], + } + ) + 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, - ) -> list[ChatModel]: - with get_db_context(db) as db: - query = db.query(Chat).filter_by(user_id=user_id) + db: Optional[AsyncSession] = None, + ) -> list[ChatTitleIdResponse]: + 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}') - all_chats = query.all() - log.debug(f'all_chats: {all_chats}') - return [ChatModel.model_validate(chat) for chat in all_chats] + stmt = stmt.order_by(Chat.updated_at.desc(), Chat.id) - 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 + if skip: + stmt = stmt.offset(skip) + if limit: + stmt = stmt.limit(limit) + + result = await db.execute(stmt) + all_chats = result.all() + return [ + ChatTitleIdResponse.model_validate( + { + 'id': chat[0], + 'title': chat[1], + 'updated_at': chat[2], + 'created_at': chat[3], + 'last_read_at': chat[4], + } + ) + for chat in all_chats + ] + + 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 @@ -1348,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() @@ -1380,122 +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(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(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(ChatMessage).filter(ChatMessage.chat_id.in_(chat_id_subquery)).delete( - 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(Chat).filter_by(user_id=user_id).delete() - 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: - 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(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, folder_id=folder_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 move_chats_by_user_id_and_folder_id( + async def delete_chats_by_user_id_and_folder_id( + self, user_id: str, folder_id: str, db: Optional[AsyncSession] = None + ) -> bool: + try: + 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) + ) + 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 + + 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])) @@ -1503,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 = [ @@ -1522,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 aa6c7bdcae..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,92 +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: + 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: @@ -262,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, @@ -278,25 +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_feedbacks_for_leaderboard(self, db: Optional[Session] = None) -> list[LeaderboardFeedbackData]: + async def get_distinct_model_ids(self, db: Optional[AsyncSession] = None) -> list[str]: + """Get distinct model_ids from feedback data for filter dropdowns.""" + 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() + ) + rows = result.all() + return sorted([row[0] for row in rows if row[0]]) + + 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. @@ -306,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 @@ -358,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 @@ -389,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 @@ -413,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 9d2b5819bc..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,67 +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'}) - result = db.query(Model).filter_by(id=id).update(data) + data['updated_at'] = int(time.time()) + 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 delete_model_by_id(self, id: str, db: Optional[Session] = None) -> bool: + 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: - 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: + result = await db.execute(select(Model).filter_by(id=id)) + model_obj = result.scalars().first() + if not model_obj: + return None + 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 + + async def delete_model_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool: + try: + 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 @@ -462,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( @@ -478,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 8e14387a78..69744e7219 100644 --- a/backend/open_webui/routers/audio.py +++ b/backend/open_webui/routers/audio.py @@ -168,6 +168,8 @@ class TTSConfigForm(BaseModel): AZURE_SPEECH_REGION: str AZURE_SPEECH_BASE_URL: str AZURE_SPEECH_OUTPUT_FORMAT: str + MISTRAL_API_KEY: str + MISTRAL_API_BASE_URL: str class STTConfigForm(BaseModel): @@ -208,6 +210,8 @@ async def get_audio_config(request: Request, user=Depends(get_admin_user)): 'AZURE_SPEECH_REGION': request.app.state.config.TTS_AZURE_SPEECH_REGION, 'AZURE_SPEECH_BASE_URL': request.app.state.config.TTS_AZURE_SPEECH_BASE_URL, 'AZURE_SPEECH_OUTPUT_FORMAT': request.app.state.config.TTS_AZURE_SPEECH_OUTPUT_FORMAT, + 'MISTRAL_API_KEY': request.app.state.config.TTS_MISTRAL_API_KEY, + 'MISTRAL_API_BASE_URL': request.app.state.config.TTS_MISTRAL_API_BASE_URL, }, 'stt': { 'OPENAI_API_BASE_URL': request.app.state.config.STT_OPENAI_API_BASE_URL, @@ -242,6 +246,8 @@ async def update_audio_config(request: Request, form_data: AudioConfigUpdateForm request.app.state.config.TTS_AZURE_SPEECH_REGION = form_data.tts.AZURE_SPEECH_REGION request.app.state.config.TTS_AZURE_SPEECH_BASE_URL = form_data.tts.AZURE_SPEECH_BASE_URL request.app.state.config.TTS_AZURE_SPEECH_OUTPUT_FORMAT = form_data.tts.AZURE_SPEECH_OUTPUT_FORMAT + request.app.state.config.TTS_MISTRAL_API_KEY = form_data.tts.MISTRAL_API_KEY + request.app.state.config.TTS_MISTRAL_API_BASE_URL = form_data.tts.MISTRAL_API_BASE_URL request.app.state.config.STT_OPENAI_API_BASE_URL = form_data.stt.OPENAI_API_BASE_URL request.app.state.config.STT_OPENAI_API_KEY = form_data.stt.OPENAI_API_KEY @@ -280,6 +286,8 @@ async def update_audio_config(request: Request, form_data: AudioConfigUpdateForm 'AZURE_SPEECH_REGION': request.app.state.config.TTS_AZURE_SPEECH_REGION, 'AZURE_SPEECH_BASE_URL': request.app.state.config.TTS_AZURE_SPEECH_BASE_URL, 'AZURE_SPEECH_OUTPUT_FORMAT': request.app.state.config.TTS_AZURE_SPEECH_OUTPUT_FORMAT, + 'MISTRAL_API_KEY': request.app.state.config.TTS_MISTRAL_API_KEY, + 'MISTRAL_API_BASE_URL': request.app.state.config.TTS_MISTRAL_API_BASE_URL, }, 'stt': { 'OPENAI_API_BASE_URL': request.app.state.config.STT_OPENAI_API_BASE_URL, @@ -322,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, @@ -551,6 +561,77 @@ async def speech(request: Request, user=Depends(get_verified_user)): return FileResponse(file_path) + elif request.app.state.config.TTS_ENGINE == 'mistral': + api_key = request.app.state.config.TTS_MISTRAL_API_KEY + api_base_url = request.app.state.config.TTS_MISTRAL_API_BASE_URL or 'https://api.mistral.ai/v1' + + if not api_key: + raise HTTPException( + status_code=400, + detail='Mistral API key is required for Mistral TTS', + ) + + try: + timeout = aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT) + async with aiohttp.ClientSession(timeout=timeout, trust_env=True) as session: + mistral_payload = { + 'input': payload.get('input', ''), + 'model': request.app.state.config.TTS_MODEL or 'mistral-tts-latest', + 'voice_id': payload.get('voice', ''), + 'response_format': 'mp3', + } + + r = await session.post( + url=f'{api_base_url}/audio/speech', + json=mistral_payload, + headers={ + 'Content-Type': 'application/json', + 'Authorization': f'Bearer {api_key}', + }, + ssl=AIOHTTP_CLIENT_SESSION_SSL, + ) + + r.raise_for_status() + + res = await r.json() + audio_data = res.get('audio_data', '') + if not audio_data: + raise ValueError('No audio_data in Mistral TTS response') + + audio_bytes = base64.b64decode(audio_data) + + async with aiofiles.open(file_path, 'wb') as f: + await f.write(audio_bytes) + + async with aiofiles.open(file_body_path, 'w') as f: + await f.write(json.dumps(payload)) + + return FileResponse(file_path) + + except Exception as e: + log.exception(e) + detail = None + + status_code = 500 + detail = 'Open WebUI: Server Connection Error' + + if r is not None: + status_code = r.status + + try: + res = await r.json() + if 'error' in res: + detail = f'External: {res["error"]}' + elif 'message' in res: + detail = f'External: {res["message"]}' + except Exception: + detail = f'External: {e}' + + raise HTTPException( + status_code=status_code, + detail=detail, + ) + def transcription_handler(request, file_path, metadata, user=None): filename = os.path.basename(file_path) @@ -582,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) @@ -620,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) @@ -689,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) @@ -796,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) @@ -981,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) @@ -1130,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, @@ -1159,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)): @@ -1238,6 +1321,8 @@ def get_available_models(request: Request) -> list[dict]: available_models = [{'name': model['name'], 'id': model['model_id']} for model in models] except requests.RequestException as e: log.error(f'Error fetching voices: {str(e)}') + elif request.app.state.config.TTS_ENGINE == 'mistral': + available_models = [{'id': 'mistral-tts-latest'}] return available_models @@ -1301,6 +1386,29 @@ def get_available_voices(request) -> dict: available_voices[voice['ShortName']] = f'{voice["DisplayName"]} ({voice["ShortName"]})' except requests.RequestException as e: log.error(f'Error fetching voices: {str(e)}') + elif request.app.state.config.TTS_ENGINE == 'mistral': + api_key = request.app.state.config.TTS_MISTRAL_API_KEY + api_base_url = request.app.state.config.TTS_MISTRAL_API_BASE_URL or 'https://api.mistral.ai/v1' + + if api_key: + try: + response = requests.get( + f'{api_base_url}/audio/voices', + headers={ + 'Authorization': f'Bearer {api_key}', + }, + timeout=AIOHTTP_CLIENT_TIMEOUT_MODEL_LIST, + ) + response.raise_for_status() + voices_data = response.json() + + for voice in voices_data: + voice_id = voice.get('voice_id', voice.get('id', '')) + voice_name = voice.get('name', voice_id) + if voice_id: + available_voices[voice_id] = voice_name + except requests.RequestException as e: + log.error(f'Error fetching Mistral voices: {str(e)}') return available_voices 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 new file mode 100644 index 0000000000..9c532a8915 --- /dev/null +++ b/backend/open_webui/routers/automations.py @@ -0,0 +1,299 @@ +import asyncio +import logging + +from typing import Optional +from fastapi import APIRouter, Depends, HTTPException, Request, status +from sqlalchemy.ext.asyncio import AsyncSession + +from open_webui.models.automations import ( + Automations, + AutomationRuns, + AutomationForm, + AutomationModel, + AutomationResponse, + AutomationRunModel, + AutomationListResponse, +) +from open_webui.utils.automations import ( + validate_rrule, + 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_async_session +from open_webui.constants import ERROR_MESSAGES + +log = logging.getLogger(__name__) + +router = APIRouter() + +PAGE_ITEM_COUNT = 30 + + +############################ +# Helpers +############################ + + +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( + status_code=status.HTTP_403_FORBIDDEN, + detail=ERROR_MESSAGES.UNAUTHORIZED, + ) + + +def check_automation_access(automation, user): + if not automation: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=ERROR_MESSAGES.NOT_FOUND, + ) + if user.role != 'admin' and user.id != automation.user_id: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail=ERROR_MESSAGES.UNAUTHORIZED, + ) + + +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, + next_runs=next_n_runs_ns(automation.data['rrule'], tz=tz), + ) + + +############################ +# GetAutomationItems (paginated) +############################ + + +@router.get('/list') +async def get_automation_items( + request: Request, + query: Optional[str] = None, + status: Optional[str] = None, + page: Optional[int] = 1, + user=Depends(get_verified_user), + db: AsyncSession = Depends(get_async_session), +): + await check_automations_permission(request, user) + limit = PAGE_ITEM_COUNT + page = max(1, page) + skip = (page - 1) * limit + + result = await Automations.search_automations( + user_id=user.id, + query=query, + status=status, + skip=skip, + limit=limit, + 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': [ + AutomationResponse( + **item.model_dump(), + last_run=latest_runs.get(item.id), + ) + for item in result.items + ], + 'total': result.total, + } + + +############################ +# CreateNewAutomation +############################ + + +@router.post('/create', response_model=AutomationResponse) +async def create_new_automation( + request: Request, + form_data: AutomationForm, + user=Depends(get_verified_user), + db: AsyncSession = Depends(get_async_session), +): + await check_automations_permission(request, user) + try: + validate_rrule(form_data.data.rrule) + except ValueError as e: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=str(e), + ) + + await check_automation_limits(request, user, form_data.data.rrule, db, is_create=True) + + tz = user.timezone + automation = await Automations.insert(user.id, form_data, next_run_ns(form_data.data.rrule, tz=tz), db=db) + return await enrich_automation(automation, db, tz=tz) + + +############################ +# GetAutomationById +############################ + + +@router.get('/{id}', response_model=AutomationResponse) +async def get_automation_by_id( + request: Request, + id: str, + user=Depends(get_verified_user), + db: AsyncSession = Depends(get_async_session), +): + await check_automations_permission(request, user) + automation = await Automations.get_by_id(id, db=db) + check_automation_access(automation, user) + return await enrich_automation(automation, db, tz=user.timezone) + + +############################ +# UpdateAutomationById +############################ + + +@router.post('/{id}/update', response_model=AutomationResponse) +async def update_automation_by_id( + request: Request, + id: str, + form_data: AutomationForm, + user=Depends(get_verified_user), + db: AsyncSession = Depends(get_async_session), +): + await check_automations_permission(request, user) + automation = await Automations.get_by_id(id, db=db) + check_automation_access(automation, user) + + try: + validate_rrule(form_data.data.rrule) + except ValueError as e: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=str(e), + ) + + await check_automation_limits(request, user, form_data.data.rrule, db, is_create=False) + + tz = user.timezone + updated = await Automations.update_by_id(id, form_data, next_run_ns(form_data.data.rrule, tz=tz), db=db) + return await enrich_automation(updated, db, tz=tz) + + +############################ +# ToggleAutomationById +############################ + + +@router.post('/{id}/toggle', response_model=AutomationResponse) +async def toggle_automation_by_id( + request: Request, + id: str, + user=Depends(get_verified_user), + db: AsyncSession = Depends(get_async_session), +): + await check_automations_permission(request, user) + automation = await Automations.get_by_id(id, db=db) + check_automation_access(automation, user) + 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) + + +############################ +# RunAutomationById +############################ + + +@router.post('/{id}/run') +async def run_automation_by_id( + request: Request, + id: str, + user=Depends(get_verified_user), + db: AsyncSession = Depends(get_async_session), +): + 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 await enrich_automation(automation, db, tz=user.timezone) + + +############################ +# DeleteAutomationById +############################ + + +@router.delete('/{id}/delete') +async def delete_automation_by_id( + request: Request, + id: str, + user=Depends(get_verified_user), + db: AsyncSession = Depends(get_async_session), +): + await check_automations_permission(request, user) + automation = await Automations.get_by_id(id, db=db) + check_automation_access(automation, user) + await AutomationRuns.delete_by_automation(id, db=db) + return await Automations.delete(id, db=db) + + +############################ +# GetAutomationRuns +############################ + + +@router.get('/{id}/runs', response_model=list[AutomationRunModel]) +async def get_automation_runs( + request: Request, + id: str, + skip: int = 0, + limit: int = 50, + user=Depends(get_verified_user), + db: AsyncSession = Depends(get_async_session), +): + await check_automations_permission(request, user) + automation = await Automations.get_by_id(id, db=db) + check_automation_access(automation, user) + 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 68ea5ff7f8..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) + 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) - 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) - 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) + 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) - 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) - 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) - 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) - 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) - 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) - 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) - 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) - 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) - 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) - 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) - 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) - 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 eacc084b42..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,15 +643,16 @@ 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 = 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} - for chat in Chats.get_chats_by_folder_id_and_user_id(folder_id, user.id, skip=skip, limit=limit, db=db) + {'title': chat.title, 'id': chat.id, 'updated_at': chat.updated_at, 'last_read_at': chat.last_read_at} + for chat in chats ] except Exception as e: @@ -659,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) ############################ @@ -669,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] @@ -680,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)] ############################ @@ -690,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) @@ -705,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)] ############################ @@ -726,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 @@ -742,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, @@ -757,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) ############################ @@ -767,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) ############################ @@ -783,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 @@ -799,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, @@ -814,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()) @@ -848,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 @@ -863,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()) @@ -883,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( @@ -910,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( @@ -926,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, { @@ -934,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, @@ -972,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( @@ -988,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, @@ -1016,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 @@ -1055,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: @@ -1069,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()) @@ -1092,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, @@ -1103,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( @@ -1136,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 = { @@ -1150,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( @@ -1183,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: @@ -1211,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, @@ -1249,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: @@ -1280,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()) @@ -1296,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) @@ -1315,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() @@ -1329,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()) @@ -1348,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) @@ -1370,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 de97e172f3..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) @@ -291,38 +291,49 @@ async def update_config( } +@router.get('/feedbacks/models', response_model=list[str]) +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 @router.get('/feedbacks/all/export', response_model=list[FeedbackModel]) -async def export_all_feedbacks(user=Depends(get_admin_user), db: Session = Depends(get_session)): - feedbacks = Feedbacks.get_all_feedbacks(db=db) +async def export_all_feedbacks( + model_id: Optional[str] = None, + user=Depends(get_admin_user), + db: AsyncSession = Depends(get_async_session), +): + 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 @@ -334,8 +345,9 @@ async def get_feedbacks( order_by: Optional[str] = None, direction: Optional[str] = None, 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 @@ -347,8 +359,10 @@ async def get_feedbacks( filter['order_by'] = order_by if direction: filter['direction'] = direction + 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 @@ -357,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, @@ -370,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) @@ -387,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) @@ -401,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 6027545190..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}') @@ -207,8 +224,8 @@ def upload_file_handler( filename = os.path.basename(unsanitized_filename) file_extension = os.path.splitext(filename)[1] - # Remove the leading dot from the file extension - file_extension = file_extension[1:] if file_extension else '' + # Remove the leading dot from the file extension and lowercase it + file_extension = file_extension[1:].lower() if file_extension else '' if process and request.app.state.config.ALLOWED_FILE_EXTENSIONS: request.app.state.config.ALLOWED_FILE_EXTENSIONS = [ @@ -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 5a56e11b68..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,15 @@ async def create_new_model( ) else: - model = Models.insert_new_model(form_data, user.id, db=db) + form_data.access_grants = await filter_allowed_access_grants( + request.app.state.config.USER_PERMISSIONS, + user.id, + user.role, + form_data.access_grants, + 'sharing.public_models', + ) + + model = await Models.insert_new_model(form_data, user.id, db=db) if model: return model else: @@ -215,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, @@ -229,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) ############################ @@ -248,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, @@ -270,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: @@ -285,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') @@ -314,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) ########################### @@ -330,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, @@ -349,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, @@ -376,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 @@ -418,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, @@ -432,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 @@ -460,11 +468,12 @@ async def toggle_model_by_id(id: str, user=Depends(get_verified_user), db: Sessi @router.post('/model/update', response_model=Optional[ModelModel]) 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, @@ -473,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, @@ -487,7 +496,15 @@ async def update_model_by_id( detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - model = Models.update_model_by_id(form_data.id, ModelForm(**form_data.model_dump()), db=db) + form_data.access_grants = await filter_allowed_access_grants( + request.app.state.config.USER_PERMISSIONS, + user.id, + user.role, + form_data.access_grants, + 'sharing.public_models', + ) + + model = await Models.update_model_by_id(form_data.id, ModelForm(**form_data.model_dump()), db=db) return model @@ -507,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. @@ -519,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, @@ -537,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, @@ -551,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, @@ -559,9 +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) - return Models.get_model_by_id(form_data.id, db=db) + await Models.update_model_updated_at_by_id(form_data.id, db=db) + + return await Models.get_model_by_id(form_data.id, db=db) ############################ @@ -573,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, @@ -585,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, @@ -598,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 705d86e1c9..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,8 +175,17 @@ async def create_new_note( detail=ERROR_MESSAGES.UNAUTHORIZED, ) + form_data.access_grants = await filter_allowed_access_grants( + request.app.state.config.USER_PERMISSIONS, + user.id, + user.role, + form_data.access_grants, + 'sharing.public_notes', + db=db, + ) + 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) @@ -197,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( @@ -207,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, @@ -228,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, @@ -252,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( @@ -262,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, @@ -278,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, @@ -288,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(), @@ -316,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( @@ -326,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, @@ -342,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, @@ -350,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) ############################ @@ -365,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( @@ -375,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, @@ -392,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 d03c37ae1a..0c6fe73abf 100644 --- a/backend/open_webui/routers/ollama.py +++ b/backend/open_webui/routers/ollama.py @@ -15,7 +15,7 @@ from typing import Optional, Union from urllib.parse import urlparse import aiohttp from aiocache import cached -import requests + from open_webui.utils.headers import include_user_info_headers from open_webui.models.chats import Chats @@ -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 ( @@ -107,19 +111,21 @@ async def send_get_request(url, key=None, user: UserModel = None): return None -async def send_post_request( +async def send_request( url: str, - payload: Union[str, bytes], - stream: bool = True, + method: str = 'POST', + *, + payload: Optional[Union[str, bytes]] = None, key: Optional[str] = None, - content_type: Optional[str] = None, user: UserModel = None, + stream: bool = False, + content_type: Optional[str] = None, metadata: Optional[dict] = None, ): 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', @@ -131,57 +137,58 @@ async def send_post_request( if metadata and metadata.get('chat_id'): headers[FORWARD_SESSION_INFO_HEADER_CHAT_ID] = metadata.get('chat_id') - r = await session.post( + r = await session.request( + method, url, data=payload, headers=headers, ssl=AIOHTTP_CLIENT_SESSION_SSL, + timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT), ) - if r.ok is False: + if not r.ok: try: res = await r.json() - await cleanup_response(r, session) if 'error' in res: raise HTTPException(status_code=r.status, detail=res['error']) - except HTTPException as e: - raise e # Re-raise HTTPException to be handled by FastAPI + except HTTPException: + raise except Exception as e: log.error(f'Failed to parse error response: {e}') - raise HTTPException( - status_code=r.status, - detail=f'Open WebUI: Server Connection Error', - ) + raise HTTPException( + status_code=r.status, + detail='Open WebUI: Server Connection Error', + ) + + r.raise_for_status() - r.raise_for_status() # Raises an error for bad responses (4xx, 5xx) if stream: response_headers = dict(r.headers) - if content_type: response_headers['Content-Type'] = content_type streaming = True return StreamingResponse( - stream_wrapper(r, session), + stream_wrapper(r), status_code=r.status, headers=response_headers, ) else: - res = await r.json() - return res + try: + return await r.json() + except Exception: + return None - except HTTPException as e: - raise e # Re-raise HTTPException to be handled by FastAPI + except HTTPException: + raise except Exception as e: - detail = f'Ollama: {e}' - raise HTTPException( status_code=r.status if r else 500, - detail=detail if e else 'Open WebUI: Server Connection Error', + detail=f'Ollama: {e}' if str(e) else 'Open WebUI: Server Connection Error', ) finally: if not streaming: - await cleanup_response(r, session) + await cleanup_response(r) def get_api_key(idx, url, configs): @@ -395,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()), @@ -430,40 +437,7 @@ async def get_ollama_tags(request: Request, url_idx: Optional[int] = None, user= else: url = request.app.state.config.OLLAMA_BASE_URLS[url_idx] key = get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS) - - r = None - try: - headers = { - **({'Authorization': f'Bearer {key}'} if key else {}), - } - - if ENABLE_FORWARD_USER_INFO_HEADERS and user: - headers = include_user_info_headers(headers, user) - - r = requests.request( - method='GET', - url=f'{url}/api/tags', - headers=headers, - ) - r.raise_for_status() - - models = r.json() - except Exception as e: - log.exception(e) - - detail = None - if r is not None: - try: - res = r.json() - if 'error' in res: - detail = f'Ollama: {res["error"]}' - except Exception: - detail = f'Ollama: {e}' - - raise HTTPException( - status_code=r.status_code if r else 500, - detail=detail if detail else 'Open WebUI: Server Connection Error', - ) + models = await send_request(f'{url}/api/tags', 'GET', key=key, user=user) if user.role == 'user' and not BYPASS_MODEL_ACCESS_CONTROL: models['models'] = await get_filtered_models(models, user) @@ -569,29 +543,7 @@ async def get_ollama_versions(request: Request, url_idx: Optional[int] = None): ) else: url = request.app.state.config.OLLAMA_BASE_URLS[url_idx] - - r = None - try: - r = requests.request(method='GET', url=f'{url}/api/version') - r.raise_for_status() - - return r.json() - except Exception as e: - log.exception(e) - - detail = None - if r is not None: - try: - res = r.json() - if 'error' in res: - detail = f'Ollama: {res["error"]}' - except Exception: - detail = f'Ollama: {e}' - - raise HTTPException( - status_code=r.status_code if r else 500, - detail=detail if detail else 'Open WebUI: Server Connection Error', - ) + return await send_request(f'{url}/api/version', 'GET') else: return {'version': False} @@ -640,10 +592,9 @@ async def unload_model( payload = {'model': model_name, 'keep_alive': 0, 'prompt': ''} try: - res = await send_post_request( - url=f'{url}/api/generate', + res = await send_request( + f'{url}/api/generate', payload=json.dumps(payload), - stream=False, key=key, user=user, ) @@ -681,11 +632,12 @@ async def pull_model( # Admin should be able to pull models from any source payload = {**form_data, 'insecure': True} - return await send_post_request( - url=f'{url}/api/pull', + return await send_request( + f'{url}/api/pull', payload=json.dumps(payload), key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS), user=user, + stream=True, ) @@ -721,11 +673,12 @@ async def push_model( url = request.app.state.config.OLLAMA_BASE_URLS[url_idx] log.debug(f'url: {url}') - return await send_post_request( - url=f'{url}/api/push', + return await send_request( + f'{url}/api/push', payload=form_data.model_dump_json(exclude_none=True).encode(), key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS), user=user, + stream=True, ) @@ -751,11 +704,12 @@ async def create_model( log.debug(f'form_data: {form_data}') url = request.app.state.config.OLLAMA_BASE_URLS[url_idx] - return await send_post_request( - url=f'{url}/api/create', + return await send_request( + f'{url}/api/create', payload=form_data.model_dump_json(exclude_none=True).encode(), key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS), user=user, + stream=True, ) @@ -790,41 +744,13 @@ async def copy_model( url = request.app.state.config.OLLAMA_BASE_URLS[url_idx] key = get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS) - try: - headers = { - 'Content-Type': 'application/json', - **({'Authorization': f'Bearer {key}'} if key else {}), - } - - if ENABLE_FORWARD_USER_INFO_HEADERS and user: - headers = include_user_info_headers(headers, user) - - r = requests.request( - method='POST', - url=f'{url}/api/copy', - headers=headers, - data=form_data.model_dump_json(exclude_none=True).encode(), - ) - r.raise_for_status() - - log.debug(f'r.text: {r.text}') - return True - except Exception as e: - log.exception(e) - - detail = None - if r is not None: - try: - res = r.json() - if 'error' in res: - detail = f'Ollama: {res["error"]}' - except Exception: - detail = f'Ollama: {e}' - - raise HTTPException( - status_code=r.status_code if r else 500, - detail=detail if detail else 'Open WebUI: Server Connection Error', - ) + await send_request( + f'{url}/api/copy', + payload=form_data.model_dump_json(exclude_none=True).encode(), + key=key, + user=user, + ) + return True @router.delete('/api/delete') @@ -858,42 +784,14 @@ async def delete_model( url = request.app.state.config.OLLAMA_BASE_URLS[url_idx] key = get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS) - r = None - try: - headers = { - 'Content-Type': 'application/json', - **({'Authorization': f'Bearer {key}'} if key else {}), - } - - if ENABLE_FORWARD_USER_INFO_HEADERS and user: - headers = include_user_info_headers(headers, user) - - r = requests.request( - method='DELETE', - url=f'{url}/api/delete', - headers=headers, - json=form_data, - ) - r.raise_for_status() - - log.debug(f'r.text: {r.text}') - return True - except Exception as e: - log.exception(e) - - detail = None - if r is not None: - try: - res = r.json() - if 'error' in res: - detail = f'Ollama: {res["error"]}' - except Exception: - detail = f'Ollama: {e}' - - raise HTTPException( - status_code=r.status_code if r else 500, - detail=detail if detail else 'Open WebUI: Server Connection Error', - ) + await send_request( + f'{url}/api/delete', + 'DELETE', + payload=json.dumps(form_data), + key=key, + user=user, + ) + return True @router.post('/api/show') @@ -904,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, @@ -920,35 +821,12 @@ async def show_model_info(request: Request, form_data: ModelNameForm, user=Depen url = request.app.state.config.OLLAMA_BASE_URLS[url_idx] key = get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS) - try: - headers = { - 'Content-Type': 'application/json', - **({'Authorization': f'Bearer {key}'} if key else {}), - } - - if ENABLE_FORWARD_USER_INFO_HEADERS and user: - headers = include_user_info_headers(headers, user) - - r = requests.request(method='POST', url=f'{url}/api/show', headers=headers, json=form_data) - r.raise_for_status() - - return r.json() - except Exception as e: - log.exception(e) - - detail = None - if r is not None: - try: - res = r.json() - if 'error' in res: - detail = f'Ollama: {res["error"]}' - except Exception: - detail = f'Ollama: {e}' - - raise HTTPException( - status_code=r.status_code if r else 500, - detail=detail if detail else 'Open WebUI: Server Connection Error', - ) + return await send_request( + f'{url}/api/show', + payload=json.dumps(form_data), + key=key, + user=user, + ) class GenerateEmbedForm(BaseModel): @@ -976,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 @@ -1004,41 +885,12 @@ async def embed( if prefix_id: form_data.model = form_data.model.replace(f'{prefix_id}.', '') - try: - headers = { - 'Content-Type': 'application/json', - **({'Authorization': f'Bearer {key}'} if key else {}), - } - - if ENABLE_FORWARD_USER_INFO_HEADERS and user: - headers = include_user_info_headers(headers, user) - - r = requests.request( - method='POST', - url=f'{url}/api/embed', - headers=headers, - data=form_data.model_dump_json(exclude_none=True).encode(), - ) - r.raise_for_status() - - data = r.json() - return data - except Exception as e: - log.exception(e) - - detail = None - if r is not None: - try: - res = r.json() - if 'error' in res: - detail = f'Ollama: {res["error"]}' - except Exception: - detail = f'Ollama: {e}' - - raise HTTPException( - status_code=r.status_code if r else 500, - detail=detail if detail else 'Open WebUI: Server Connection Error', - ) + return await send_request( + f'{url}/api/embed', + payload=form_data.model_dump_json(exclude_none=True).encode(), + key=key, + user=user, + ) class GenerateEmbeddingsForm(BaseModel): @@ -1061,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 @@ -1089,41 +944,12 @@ async def embeddings( if prefix_id: form_data.model = form_data.model.replace(f'{prefix_id}.', '') - try: - headers = { - 'Content-Type': 'application/json', - **({'Authorization': f'Bearer {key}'} if key else {}), - } - - if ENABLE_FORWARD_USER_INFO_HEADERS and user: - headers = include_user_info_headers(headers, user) - - r = requests.request( - method='POST', - url=f'{url}/api/embeddings', - headers=headers, - data=form_data.model_dump_json(exclude_none=True).encode(), - ) - r.raise_for_status() - - data = r.json() - return data - except Exception as e: - log.exception(e) - - detail = None - if r is not None: - try: - res = r.json() - if 'error' in res: - detail = f'Ollama: {res["error"]}' - except Exception: - detail = f'Ollama: {e}' - - raise HTTPException( - status_code=r.status_code if r else 500, - detail=detail if detail else 'Open WebUI: Server Connection Error', - ) + return await send_request( + f'{url}/api/embeddings', + payload=form_data.model_dump_json(exclude_none=True).encode(), + key=key, + user=user, + ) class GenerateCompletionForm(BaseModel): @@ -1152,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: @@ -1175,11 +1005,12 @@ async def generate_completion( if prefix_id: form_data.model = form_data.model.replace(f'{prefix_id}.', '') - return await send_post_request( - url=f'{url}/api/generate', + return await send_request( + f'{url}/api/generate', payload=form_data.model_dump_json(exclude_none=True).encode(), key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS), user=user, + stream=True, ) @@ -1238,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. @@ -1267,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: @@ -1285,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( @@ -1319,13 +1130,13 @@ async def generate_chat_completion( if prefix_id: payload['model'] = payload['model'].replace(f'{prefix_id}.', '') - return await send_post_request( - url=f'{url}/api/chat', + return await send_request( + f'{url}/api/chat', payload=json.dumps(payload), - stream=form_data.stream, key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS), - content_type='application/x-ndjson', user=user, + stream=form_data.stream, + content_type='application/x-ndjson', metadata=metadata, ) @@ -1365,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. @@ -1385,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 @@ -1394,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( @@ -1429,12 +1220,12 @@ async def generate_openai_completion( if prefix_id: payload['model'] = payload['model'].replace(f'{prefix_id}.', '') - return await send_post_request( - url=f'{url}/v1/completions', + return await send_request( + f'{url}/v1/completions', payload=json.dumps(payload), - stream=payload.get('stream', False), key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS), user=user, + stream=payload.get('stream', False), metadata=metadata, ) @@ -1447,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. @@ -1467,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 @@ -1480,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( @@ -1514,12 +1285,12 @@ async def generate_openai_chat_completion( if prefix_id: payload['model'] = payload['model'].replace(f'{prefix_id}.', '') - return await send_post_request( - url=f'{url}/v1/chat/completions', + return await send_request( + f'{url}/v1/chat/completions', payload=json.dumps(payload), - stream=payload.get('stream', False), key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS), user=user, + stream=payload.get('stream', False), metadata=metadata, ) @@ -1547,17 +1318,75 @@ 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 + + await check_model_access(user, model_info) + else: + 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( + str(url_idx), + request.app.state.config.OLLAMA_API_CONFIGS.get(url, {}), # Legacy support + ) + + prefix_id = api_config.get('prefix_id', None) + if prefix_id: + payload['model'] = payload['model'].replace(f'{prefix_id}.', '') + + return await send_request( + f'{url}/v1/messages', + payload=json.dumps(payload), + key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS), + user=user, + stream=payload.get('stream', False), + content_type='text/event-stream' if payload.get('stream', False) else None, + ) + + +class ResponsesForm(BaseModel): + model: str + + model_config = ConfigDict(extra='allow') + + +@router.post('/v1/responses') +@router.post('/v1/responses/{url_idx}') +async def generate_responses( + request: Request, + form_data: ResponsesForm, + url_idx: Optional[int] = None, + user=Depends(get_verified_user), +): + """ + Proxy for Ollama's OpenAI-compatible /v1/responses endpoint. + + Forwards the request as-is to the Ollama backend, applying the same + model resolution, access control, and prefix_id handling used by + the OpenAI-compatible /v1/chat/completions proxy. + + See https://ollama.com/blog/responses-api + """ + if not request.app.state.config.ENABLE_OLLAMA_API: + raise HTTPException(status_code=503, detail='Ollama API is disabled') + + payload = form_data.model_dump() + model_id = form_data.model + + 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, @@ -1586,13 +1415,13 @@ async def generate_anthropic_messages( if prefix_id: payload['model'] = payload['model'].replace(f'{prefix_id}.', '') - return await send_post_request( - url=f'{url}/v1/messages', + return await send_request( + f'{url}/v1/responses', payload=json.dumps(payload), - stream=payload.get('stream', False), - content_type='text/event-stream' if payload.get('stream', False) else None, key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS), user=user, + stream=payload.get('stream', False), + content_type='text/event-stream' if payload.get('stream', False) else None, ) @@ -1602,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: @@ -1619,45 +1448,26 @@ async def get_openai_models( else: url = request.app.state.config.OLLAMA_BASE_URLS[url_idx] - try: - r = requests.request(method='GET', url=f'{url}/api/tags') - r.raise_for_status() + model_list = await send_request(f'{url}/api/tags', 'GET') - model_list = r.json() - - models = [ - { - 'id': model['model'], - 'object': 'model', - 'created': int(time.time()), - 'owned_by': 'openai', - } - for model in models['models'] - ] - except Exception as e: - log.exception(e) - error_detail = 'Open WebUI: Server Connection Error' - if r is not None: - try: - res = r.json() - if 'error' in res: - error_detail = f'Ollama: {res["error"]}' - except Exception: - error_detail = f'Ollama: {e}' - - raise HTTPException( - status_code=r.status_code if r else 500, - detail=error_detail, - ) + models = [ + { + 'id': model['model'], + 'object': 'model', + 'created': int(time.time()), + 'owned_by': 'openai', + } + for model in model_list.get('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()), @@ -1734,13 +1544,16 @@ async def download_file_stream(ollama_url, file_url, file_path, file_name, chunk file.close() hashed = calculate_sha256(file_path, chunk_size) - with open(file_path, 'rb') as file: - chunk_size = 1024 * 1024 * 2 - url = f'{ollama_url}/api/blobs/sha256:{hashed}' - with requests.Session() as session: - response = session.post(url, data=file, timeout=30) + with open(file_path, 'rb') as f: + blob_data = f.read() - if response.ok: + url = f'{ollama_url}/api/blobs/sha256:{hashed}' + blob_timeout = aiohttp.ClientTimeout(total=30) + async with aiohttp.ClientSession(timeout=blob_timeout, trust_env=True) as blob_session: + async with blob_session.post( + url, data=blob_data, ssl=AIOHTTP_CLIENT_SESSION_SSL + ) as blob_response: + if blob_response.ok: res = { 'done': done, 'blob': f'sha256:{hashed}', @@ -1777,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), @@ -1836,47 +1649,51 @@ async def upload_model( # --- P3: Upload to ollama /api/blobs --- with open(file_path, 'rb') as f: - url = f'{ollama_url}/api/blobs/sha256:{file_hash}' - response = requests.post(url, data=f) + blob_data = f.read() - if response.ok: - log.info(f'Uploaded to /api/blobs') # DEBUG - # Remove local file - os.remove(file_path) + 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: + if not response.ok: + raise Exception('Ollama: Could not create blob, Please try again.') - # Create model in ollama - model_name, ext = os.path.splitext(filename) - log.info(f'Created Model: {model_name}') # DEBUG + log.info(f'Uploaded to /api/blobs') # DEBUG + # Remove local file + os.remove(file_path) - create_payload = { - 'model': model_name, - # Reference the file by its original name => the uploaded blob's digest - 'files': {filename: f'sha256:{file_hash}'}, - } - log.info(f'Model Payload: {create_payload}') # DEBUG + # Create model in ollama + model_name, ext = os.path.splitext(filename) + log.info(f'Created Model: {model_name}') # DEBUG - # Call ollama /api/create - # https://github.com/ollama/ollama/blob/main/docs/api.md#create-a-model - create_resp = requests.post( - url=f'{ollama_url}/api/create', + create_payload = { + 'model': model_name, + # Reference the file by its original name => the uploaded blob's digest + 'files': {filename: f'sha256:{file_hash}'}, + } + log.info(f'Model Payload: {create_payload}') # DEBUG + + # Call ollama /api/create + # https://github.com/ollama/ollama/blob/main/docs/api.md#create-a-model + async with aiohttp.ClientSession(timeout=upload_timeout, trust_env=True) as create_session: + async with create_session.post( + f'{ollama_url}/api/create', headers={'Content-Type': 'application/json'}, data=json.dumps(create_payload), - ) - - if create_resp.ok: - log.info(f'API SUCCESS!') # DEBUG - done_msg = { - 'done': True, - 'blob': f'sha256:{file_hash}', - 'name': filename, - 'model_created': model_name, - } - yield f'data: {json.dumps(done_msg)}\n\n' - else: - raise Exception(f'Failed to create model in Ollama. {create_resp.text}') - - else: - raise Exception('Ollama: Could not create blob, Please try again.') + ssl=AIOHTTP_CLIENT_SESSION_SSL, + ) as create_resp: + if create_resp.ok: + log.info(f'API SUCCESS!') # DEBUG + done_msg = { + 'done': True, + 'blob': f'sha256:{file_hash}', + 'name': filename, + 'model_created': model_name, + } + yield f'data: {json.dumps(done_msg)}\n\n' + else: + resp_text = await create_resp.text() + raise Exception(f'Failed to create model in Ollama. {resp_text}') except Exception as e: res = {'error': str(e)} diff --git a/backend/open_webui/routers/openai.py b/backend/open_webui/routers/openai.py index 0e7c67c1f6..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 @@ -1131,21 +1134,31 @@ async def generate_chat_completion( is_responses = api_config.get('api_type') == 'responses' if api_config.get('azure', False): - api_version = api_config.get('api_version', '2023-03-15-preview') - request_url, payload = convert_to_azure_payload(url, payload, api_version) - # Only set api-key header if not using Azure Entra ID authentication auth_type = api_config.get('auth_type', 'bearer') if auth_type not in ('azure_ad', 'microsoft_entra_id'): headers['api-key'] = key - headers['api-version'] = api_version + # Azure v1 format: base URL already ends with /openai/v1, + # model stays in the payload, no deployment URL rewriting. + is_azure_v1 = bool(re.search(r'/openai/v1(?:/|$)', url)) - if is_responses: - payload = convert_to_responses_payload(payload) - request_url = f'{request_url}/responses?api-version={api_version}' + if is_azure_v1: + if is_responses: + payload = convert_to_responses_payload(payload) + request_url = f'{url.rstrip("/")}/responses' + else: + request_url = f'{url.rstrip("/")}/chat/completions' else: - request_url = f'{request_url}/chat/completions?api-version={api_version}' + api_version = api_config.get('api_version', '2023-03-15-preview') + request_url, payload = convert_to_azure_payload(url, payload, api_version) + headers['api-version'] = api_version + + if is_responses: + payload = convert_to_responses_payload(payload) + request_url = f'{request_url}/responses?api-version={api_version}' + else: + request_url = f'{request_url}/chat/completions?api-version={api_version}' else: if is_responses: payload = convert_to_responses_payload(payload) @@ -1164,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', @@ -1178,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), ) @@ -1215,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): @@ -1251,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), ) @@ -1296,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): @@ -1330,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: @@ -1350,30 +1384,29 @@ async def responses( ) r = None - session = None streaming = False try: headers, cookies = await get_headers_and_cookies(request, url, key, api_config, user=user) if api_config.get('azure', False): - api_version = api_config.get('api_version', '2023-03-15-preview') - auth_type = api_config.get('auth_type', 'bearer') if auth_type not in ('azure_ad', 'microsoft_entra_id'): headers['api-key'] = key - headers['api-version'] = api_version + is_azure_v1 = bool(re.search(r'/openai/v1(?:/|$)', url)) - model = payload.get('model', '') - request_url = f'{url}/openai/deployments/{model}/responses?api-version={api_version}' + if is_azure_v1: + request_url = f'{url.rstrip("/")}/responses' + else: + api_version = api_config.get('api_version', '2023-03-15-preview') + headers['api-version'] = api_version + 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, @@ -1381,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), ) @@ -1405,6 +1439,8 @@ async def responses( return response_data + except HTTPException: + raise except Exception as e: log.exception(e) raise HTTPException( @@ -1413,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 @@ -1452,34 +1495,35 @@ async def proxy(path: str, request: Request, user=Depends(get_verified_user)): ) r = None - session = None streaming = False try: headers, cookies = await get_headers_and_cookies(request, url, key, api_config, user=user) if api_config.get('azure', False): - api_version = api_config.get('api_version', '2023-03-15-preview') - # Only set api-key header if not using Azure Entra ID authentication auth_type = api_config.get('auth_type', 'bearer') if auth_type not in ('azure_ad', 'microsoft_entra_id'): headers['api-key'] = key - headers['api-version'] = api_version + is_azure_v1 = bool(re.search(r'/openai/v1(?:/|$)', url)) - payload = json.loads(body) - url, payload = convert_to_azure_payload(url, payload, api_version) - body = json.dumps(payload).encode() + if is_azure_v1: + qs = request.url.query + request_url = f'{url.rstrip("/")}/{path}' + (f'?{qs}' if qs else '') + else: + api_version = api_config.get('api_version', '2023-03-15-preview') + headers['api-version'] = api_version - request_url = f'{url}/{path}?api-version={api_version}' + payload = json.loads(body) + url, payload = convert_to_azure_payload(url, payload, api_version) + body = json.dumps(payload).encode() + + request_url = f'{url}/{path}?api-version={api_version}' 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, @@ -1487,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), ) @@ -1511,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( @@ -1519,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 e4af8bb513..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,9 +168,17 @@ async def create_new_prompt( detail=ERROR_MESSAGES.UNAUTHORIZED, ) - prompt = Prompts.get_prompt_by_command(form_data.command, db=db) + form_data.access_grants = await filter_allowed_access_grants( + request.app.state.config.USER_PERMISSIONS, + user.id, + user.role, + form_data.access_grants, + 'sharing.public_prompts', + ) + + 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 @@ -190,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, @@ -210,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, @@ -232,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, @@ -252,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, @@ -275,12 +287,13 @@ async def get_prompt_by_id(prompt_id: str, user=Depends(get_verified_user), db: @router.post('/id/{prompt_id}/update', response_model=Optional[PromptModel]) async def update_prompt_by_id( + request: Request, 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( @@ -291,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, @@ -307,15 +320,23 @@ 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 = await filter_allowed_access_grants( + request.app.state.config.USER_PERMISSIONS, + user.id, + user.role, + form_data.access_grants, + 'sharing.public_prompts', + ) + # 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: @@ -335,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( @@ -348,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, @@ -364,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: @@ -386,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, @@ -397,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, @@ -411,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: @@ -436,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, @@ -447,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, @@ -461,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, @@ -469,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) ############################ @@ -480,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( @@ -491,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, @@ -505,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( @@ -520,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( @@ -531,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, @@ -545,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 @@ -559,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( @@ -576,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, @@ -589,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 @@ -598,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( @@ -613,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, @@ -626,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, @@ -641,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( @@ -656,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, @@ -676,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, @@ -692,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( @@ -707,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, @@ -720,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 6c9e988dd6..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 @@ -579,9 +579,9 @@ class WebConfig(BaseModel): WEB_SEARCH_TRUST_ENV: Optional[bool] = None WEB_SEARCH_RESULT_COUNT: Optional[int] = None WEB_SEARCH_CONCURRENT_REQUESTS: Optional[int] = None + WEB_SEARCH_DOMAIN_FILTER_LIST: Optional[List[str]] = [] WEB_FETCH_MAX_CONTENT_LENGTH: Optional[int] = None WEB_LOADER_CONCURRENT_REQUESTS: Optional[int] = None - WEB_SEARCH_DOMAIN_FILTER_LIST: Optional[List[str]] = [] BYPASS_WEB_SEARCH_EMBEDDING_AND_RETRIEVAL: Optional[bool] = None BYPASS_WEB_SEARCH_WEB_LOADER: Optional[bool] = None OLLAMA_CLOUD_WEB_SEARCH_API_KEY: Optional[str] = None @@ -1174,7 +1174,7 @@ async def update_rag_config(request: Request, form_data: ConfigForm, user=Depend 'WEB_SEARCH_TRUST_ENV': request.app.state.config.WEB_SEARCH_TRUST_ENV, 'WEB_SEARCH_RESULT_COUNT': request.app.state.config.WEB_SEARCH_RESULT_COUNT, 'WEB_SEARCH_CONCURRENT_REQUESTS': request.app.state.config.WEB_SEARCH_CONCURRENT_REQUESTS, - 'FETCH_URL_MAX_CONTENT_LENGTH': request.app.state.config.FETCH_URL_MAX_CONTENT_LENGTH, + 'WEB_FETCH_MAX_CONTENT_LENGTH': request.app.state.config.WEB_FETCH_MAX_CONTENT_LENGTH, 'WEB_LOADER_CONCURRENT_REQUESTS': request.app.state.config.WEB_LOADER_CONCURRENT_REQUESTS, 'WEB_SEARCH_DOMAIN_FILTER_LIST': request.app.state.config.WEB_SEARCH_DOMAIN_FILTER_LIST, 'BYPASS_WEB_SEARCH_EMBEDDING_AND_RETRIEVAL': request.app.state.config.BYPASS_WEB_SEARCH_EMBEDDING_AND_RETRIEVAL, @@ -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 @@ -2580,11 +2580,12 @@ async def process_files_batch( request: Request, form_data: BatchProcessFilesForm, user=Depends(get_verified_user), + db=None, ) -> BatchProcessFilesResponse: """ 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. @@ -2602,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_file = await Files.get_file_by_id(file.id, db=db) if not db_file: file_errors.append( BatchProcessFilesResult( @@ -2664,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) + 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 a0b8bccd44..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( @@ -120,7 +120,7 @@ async def get_tools( auth_type = server.get('auth_type', 'none') session_token = None - if auth_type == 'oauth_2.1': + if auth_type in ('oauth_2.1', 'oauth_2.1_static'): splits = server_id.split(':') server_id = splits[-1] if len(splits) > 1 else server_id @@ -148,7 +148,7 @@ async def get_tools( { 'authenticated': session_token is not None, } - if auth_type == 'oauth_2.1' + if auth_type in ('oauth_2.1', 'oauth_2.1_static') else {} ), } @@ -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,18 +350,26 @@ 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 = await filter_allowed_access_grants( + request.app.state.config.USER_PERMISSIONS, + user.id, + user.role, + form_data.access_grants, + 'sharing.public_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) @@ -393,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, @@ -413,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, @@ -445,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, @@ -457,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, @@ -473,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 @@ -481,13 +488,21 @@ async def update_tools_by_id( specs = get_tool_specs(TOOLS[id]) + form_data.access_grants = await filter_allowed_access_grants( + request.app.state.config.USER_PERMISSIONS, + user.id, + user.role, + form_data.access_grants, + 'sharing.public_tools', + ) + updated = { **form_data.model_dump(exclude={'id'}), 'specs': specs, } 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 @@ -519,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, @@ -530,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, @@ -544,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, @@ -552,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) ############################ @@ -567,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, @@ -578,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, @@ -592,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: @@ -607,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, @@ -617,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, @@ -632,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( @@ -651,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, @@ -662,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, @@ -679,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'): @@ -702,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, @@ -713,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, @@ -730,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'): @@ -744,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}') @@ -760,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, @@ -770,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, @@ -785,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( @@ -799,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, @@ -810,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, @@ -827,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'): @@ -845,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, @@ -856,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, @@ -873,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'): @@ -883,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 b263140878..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 @@ -232,6 +233,7 @@ class FeaturesPermissions(BaseModel): image_generation: bool = True code_interpreter: bool = True memories: bool = True + automations: bool = False class SettingsPermissions(BaseModel): @@ -271,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: @@ -292,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') @@ -300,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, @@ -309,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: @@ -328,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: @@ -355,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( @@ -379,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: @@ -397,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: @@ -434,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: @@ -448,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: @@ -466,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: @@ -485,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: @@ -502,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 @@ -541,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), } @@ -558,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: @@ -572,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, @@ -586,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, @@ -604,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( @@ -638,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, @@ -655,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( @@ -678,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 33c9ffea05..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,20 @@ 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') +async def chat_events(sid, data): + user = SESSION_POOL.get(sid) + if not user: + return + + event_data = data.get('data', {}) + event_type = event_data.get('type') + + if event_type == 'last_read_at': + await Chats.update_chat_last_read_at_by_id(data['chat_id'], user['id']) def normalize_document_id(document_id: str) -> str: @@ -516,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 @@ -524,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, @@ -589,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 @@ -597,17 +628,17 @@ 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, - permission='read', + permission='write', ) ): - log.error(f'User {user.get("id")} does not have access to note {note_id}') + 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') @@ -780,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'] @@ -800,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'], ) @@ -818,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'], { @@ -830,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'], { @@ -840,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'], ) @@ -849,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'], { @@ -859,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'], ) @@ -868,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'], { @@ -880,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'], ) @@ -889,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'], { @@ -904,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 088319f2e6..61cd5ede4e 100644 --- a/backend/open_webui/tools/builtin.py +++ b/backend/open_webui/tools/builtin.py @@ -149,7 +149,7 @@ async def calculate_timestamp( async def search_web( query: str, - count: int = 5, + count: Optional[int] = None, __request__: Request = None, __user__: dict = None, ) -> str: @@ -158,7 +158,7 @@ async def search_web( or topics not covered in internal documents. :param query: The search query to look up - :param count: Number of results to return (default: 5) + :param count: Number of results to return (default: admin-configured value) :return: JSON with search results containing title, link, and snippet for each result """ if __request__ is None: @@ -168,12 +168,9 @@ async def search_web( engine = __request__.app.state.config.WEB_SEARCH_ENGINE user = UserModel(**__user__) if __user__ else None - # Enforce maximum result count from config to prevent abuse - count = ( - count - if count < __request__.app.state.config.WEB_SEARCH_RESULT_COUNT - else __request__.app.state.config.WEB_SEARCH_RESULT_COUNT - ) + configured = __request__.app.state.config.WEB_SEARCH_RESULT_COUNT + max_count = 5 if configured is None else configured + count = max(1, min(count, max_count)) if count is not None else max_count results = await asyncio.to_thread(_search_web, __request__, engine, query, user) @@ -253,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, @@ -320,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, @@ -476,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 {}, @@ -498,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 {}, @@ -653,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]) @@ -683,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 = [ @@ -733,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, @@ -763,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 '') @@ -811,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, @@ -881,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'}) @@ -924,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, @@ -950,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'}) @@ -1001,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, @@ -1076,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'}) @@ -1148,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() @@ -1204,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} @@ -1216,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, @@ -1277,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: @@ -1339,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 = [] @@ -1430,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': '', @@ -1445,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( @@ -1489,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, @@ -1504,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( @@ -1555,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__: @@ -1580,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, @@ -1597,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}, @@ -1620,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( { @@ -1636,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}, @@ -1644,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, @@ -1722,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'}) @@ -1732,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__), @@ -1814,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 @@ -1829,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, @@ -1906,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 = [] @@ -1917,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, @@ -1929,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 = { @@ -1946,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( { @@ -1957,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, @@ -2039,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: @@ -2056,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, @@ -2072,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, @@ -2102,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, @@ -2117,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': '', @@ -2196,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 @@ -2206,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, @@ -2250,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( { @@ -2274,15 +2283,15 @@ async def query_knowledge_bases( async def view_skill( - name: str, + id: str, __request__: Request = None, __user__: dict = None, ) -> str: """ - Load the full instructions of a skill by its name from the available skills manifest. + Load the full instructions of a skill by its id from the available skills manifest. Use this when you need detailed instructions for a skill listed in . - :param name: The name of the skill to load (as shown in the manifest) + :param id: The id of the skill to load (as shown in the manifest) :return: The full skill instructions as markdown content """ if __request__ is None: @@ -2297,17 +2306,17 @@ async def view_skill( user_id = __user__.get('id') - # Direct DB lookup by unique name - skill = Skills.get_skill_by_name(name) + # Direct DB lookup by id (case-insensitive since IDs are stored lowercase) + skill = await Skills.get_skill_by_id(id.lower()) if not skill or not skill.is_active: - return json.dumps({'error': f"Skill '{name}' not found"}) + return json.dumps({'error': f"Skill '{id}' not found"}) # 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, @@ -2339,14 +2348,41 @@ 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.") - status: Literal['pending', 'in_progress', 'completed', 'cancelled'] = Field('pending', description="Task status.") + id: Optional[str] = Field(None, description='Unique identifier for the task. Auto-generated if omitted.') + 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, @@ -2354,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 34412d6041..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, @@ -206,7 +205,7 @@ def create_token(data: dict, expires_delta: Union[timedelta, None] = None) -> st payload.update({'exp': expire}) jti = str(uuid.uuid4()) - payload.update({'jti': jti}) + payload.update({'jti': jti, 'iat': datetime.now(UTC)}) encoded_jwt = jwt.encode(payload, SESSION_SECRET, algorithm=ALGORITHM) return encoded_jwt @@ -221,15 +220,34 @@ def decode_token(token: str) -> Optional[dict]: async def is_valid_token(request, decoded) -> bool: - # Require Redis to check revoked tokens + """ + Check whether a JWT has been revoked. Two mechanisms: + 1. Per-token (jti) — used by user-initiated sign-out (known jti). + 2. Per-user (revoked_at) — used by OIDC back-channel logout when + individual jti values are unknown; rejects tokens with iat <= revoked_at. + """ if request.app.state.redis: + # Per-token revocation jti = decoded.get('jti') - if jti: revoked = await request.app.state.redis.get(f'{REDIS_KEY_PREFIX}:auth:token:{jti}:revoked') if revoked: return False + # 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') + if revoked_at: + try: + revoked_at_ts = int(revoked_at) + token_iat = decoded.get('iat') + # No iat means legacy token — reject since we can't verify issue time + if token_iat is None or token_iat <= revoked_at_ts: + return False + except (ValueError, TypeError): + pass + return True @@ -303,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 @@ -332,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, @@ -348,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( @@ -380,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( @@ -392,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, @@ -400,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 @@ -430,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. @@ -440,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 new file mode 100644 index 0000000000..5997b0e58b --- /dev/null +++ b/backend/open_webui/utils/automations.py @@ -0,0 +1,432 @@ +""" +Automation utilities. + +RRULE helpers, worker loop, and execution logic. +Follows the utils/.py pattern (cf. utils/channels.py, utils/task.py). + +Environment: + AUTOMATION_POLL_INTERVAL – seconds between polls (default: 10) +""" + +import asyncio +import logging +import os +import random +import time +from datetime import datetime +from typing import Optional +from uuid import uuid4 +from zoneinfo import ZoneInfo + +from dateutil.rrule import rrulestr +from fastapi import Request +from starlette.datastructures import Headers + +from open_webui.models.automations import Automations, AutomationRuns, AutomationModel +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_async_db + +log = logging.getLogger(__name__) + +AUTOMATION_POLL_INTERVAL = int(os.getenv('AUTOMATION_POLL_INTERVAL', '10')) + + +#################### +# RRULE Helpers +#################### + + +def _parse_rule(s: str): + """Parse RRULE with clock-aligned DTSTART for sub-daily frequencies. + + MINUTELY/HOURLY rules use a fixed epoch DTSTART (2000-01-01 00:00) + so intervals snap to clock boundaries (e.g. every 5min = :00, :05, :10). + """ + raw = s.replace('RRULE:', '') + parts = dict(p.split('=', 1) for p in raw.split(';') if '=' in p) + freq = parts.get('FREQ', '') + + if freq in ('MINUTELY', 'HOURLY'): + epoch = datetime(2000, 1, 1, 0, 0, 0) + return rrulestr(s, dtstart=epoch, ignoretz=True) + return rrulestr(s, ignoretz=True) + + +def validate_rrule(s: str) -> None: + """Raise ValueError if the RRULE is malformed or exhausted.""" + try: + rule = _parse_rule(s) + except Exception as e: + raise ValueError(f'Invalid RRULE: {e}') + if rule.after(datetime.now()) is None: + raise ValueError('RRULE has no future occurrences') + + +def next_run_ns(s: str, tz: str = None) -> Optional[int]: + """Next occurrence as epoch nanoseconds, respecting user timezone.""" + now = datetime.now(ZoneInfo(tz)) if tz else datetime.now() + dt = _parse_rule(s).after(now.replace(tzinfo=None)) + if dt is None: + return None + if tz: + dt = dt.replace(tzinfo=ZoneInfo(tz)) + return int(dt.timestamp() * 1_000_000_000) + + +def next_n_runs_ns(s: str, n: int = 5, tz: str = None) -> list[int]: + """Compute next N occurrences for UI preview.""" + rule = _parse_rule(s) + result = [] + dt = datetime.now() + for _ in range(n): + dt = rule.after(dt) + if not dt: + break + if tz: + dt_tz = dt.replace(tzinfo=ZoneInfo(tz)) + result.append(int(dt_tz.timestamp() * 1_000_000_000)) + else: + result.append(int(dt.timestamp() * 1_000_000_000)) + 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 +############################ + + +async def automation_worker_loop(app) -> None: + """Poll for due automations, claim, fire-and-forget execute. + + Runs on every instance. Poll interval is configurable via + AUTOMATION_POLL_INTERVAL env var (default: 10 seconds). + """ + log.info(f'Automation worker started (poll interval: {AUTOMATION_POLL_INTERVAL}s)') + while True: + try: + 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: + asyncio.create_task(execute_automation(app, automation)) + except Exception: + log.exception('Automation worker error') + + # Jitter to spread load across instances + await asyncio.sleep(AUTOMATION_POLL_INTERVAL + random.uniform(0, 2)) + + +########################## +# Execute +#################### + + +def _build_request(app) -> Request: + """Build a minimal ASGI Request for chat_completion. + + Mirrors the mock-request pattern used in main.py lifespan + (model pre-fetch, tool server init) for consistency. + """ + scope = { + 'type': 'http', + 'asgi': {'version': '3.0', 'spec_version': '2.0'}, + 'method': 'POST', + 'path': '/api/v1/automations/internal', + 'query_string': b'', + 'headers': Headers({}).raw, + 'client': ('127.0.0.1', 0), + 'server': ('127.0.0.1', 80), + 'scheme': 'http', + 'app': app, + } + request = Request(scope) + # Ensure request.state is initialized with required attributes + request.state.token = None + request.state.enable_api_keys = False + return request + + +def _resolve_model_tool_ids(app, model_id: str) -> list[str]: + """Read model-attached tool_ids from model config. + + The frontend does this in Chat.svelte (model.info.meta.toolIds). + The backend never auto-resolves them, so we must do it explicitly. + """ + models = getattr(app.state, 'MODELS', {}) + model = models.get(model_id, {}) + tool_ids = model.get('info', {}).get('meta', {}).get('toolIds', []) + return list(tool_ids) if tool_ids else [] + + +def _resolve_model_features(app, model_id: str) -> dict: + """Read model default features from model config. + + The frontend does this in Chat.svelte (model.info.meta.defaultFeatureIds + + model.info.meta.capabilities). Enables features like web_search, + code_interpreter, image_generation when the model has them as defaults + AND the capability is enabled AND the admin has enabled the feature. + """ + models = getattr(app.state, 'MODELS', {}) + model = models.get(model_id, {}) + meta = model.get('info', {}).get('meta', {}) + + default_feature_ids = meta.get('defaultFeatureIds', []) + if not default_feature_ids: + return {} + + capabilities = meta.get('capabilities', {}) + config = app.state.config + features = {} + + # code_interpreter is excluded: it requires the frontend event emitter + # and does not work in headless backend execution. + feature_checks = { + 'web_search': getattr(config, 'ENABLE_WEB_SEARCH', False), + 'image_generation': getattr(config, 'ENABLE_IMAGE_GENERATION', False), + } + + for feature_id in default_feature_ids: + if feature_id in feature_checks: + # Feature must be: in defaultFeatureIds + capability enabled + admin enabled + if capabilities.get(feature_id) and feature_checks[feature_id]: + features[feature_id] = True + + return features + + +def _resolve_model_filter_ids(app, model_id: str) -> list[str]: + """Read model default filter_ids from model config.""" + models = getattr(app.state, 'MODELS', {}) + model = models.get(model_id, {}) + filter_ids = model.get('info', {}).get('meta', {}).get('defaultFilterIds', []) + return list(filter_ids) if filter_ids else [] + + +def _resolve_model_terminal_id(app, model_id: str) -> Optional[str]: + """Read model default terminal_id from model config. + + The frontend does this in Chat.svelte (model.info.meta.terminalId). + """ + models = getattr(app.state, 'MODELS', {}) + model = models.get(model_id, {}) + return model.get('info', {}).get('meta', {}).get('terminalId') or None + + +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. + + Routes through the open-webui terminal proxy endpoint so that + auth headers, orchestrator policy routing, and X-User-Id are + handled correctly — same path the frontend uses. + """ + import aiohttp + + connections = getattr(getattr(app, 'state', None), 'config', None) + if connections is None: + return + connections = getattr(connections, 'TERMINAL_SERVER_CONNECTIONS', None) or [] + connection = next((c for c in connections if c.get('id') == server_id), None) + if connection is None: + log.warning(f'Terminal server {server_id} not found for CWD set') + return + + base_url = (connection.get('url') or '').rstrip('/') + if not base_url: + return + + # Build target URL — route through orchestrator policy if configured + policy_id = connection.get('policy_id') + if connection.get('server_type') == 'orchestrator' and policy_id: + target_url = f'{base_url}/p/{policy_id}/files/cwd' + else: + target_url = f'{base_url}/files/cwd' + + headers = {'Content-Type': 'application/json', 'X-User-Id': user.id} + if chat_id: + headers['X-Session-Id'] = chat_id + + auth_type = connection.get('auth_type', 'bearer') + if auth_type == 'bearer': + headers['Authorization'] = f'Bearer {connection.get("key", "")}' + + try: + async with aiohttp.ClientSession(timeout=aiohttp.ClientTimeout(total=10)) as session: + async with session.post( + target_url, + json={'path': cwd}, + headers=headers, + ) as resp: + if resp.status != 200: + body = await resp.text() + log.warning(f'Failed to set terminal CWD to {cwd}: HTTP {resp.status} — {body[:200]}') + except Exception as e: + log.warning(f'Failed to set terminal CWD: {e}') + + +async def execute_automation(app, automation: AutomationModel) -> None: + """Execute an automation through the full chat completion pipeline. + + Creates a real chat, then calls chat_completion exactly like the frontend: + session_id + chat_id + message_id → async task → pipeline handles everything + (filters, model params, knowledge/RAG, tools, DB saves, webhooks). + """ + try: + user = await Users.get_user_by_id(automation.user_id) + if not user: + await _record_run(automation.id, 'error', error='User not found') + return + + prompt = prompt_template(automation.data['prompt'], user) + model_id = automation.data['model_id'] + terminal_config = automation.data.get('terminal') + + # Generate proper UUIDs for messages (same as frontend) + user_msg_id = str(uuid4()) + assistant_msg_id = str(uuid4()) + + # Create the chat with user message (same structure as frontend) + chat = await Chats.insert_new_chat( + automation.user_id, + ChatForm( + chat={ + 'title': automation.name, + 'models': [model_id], + 'history': { + 'currentId': assistant_msg_id, + 'messages': { + user_msg_id: { + 'id': user_msg_id, + 'parentId': None, + 'role': 'user', + 'content': prompt, + 'childrenIds': [assistant_msg_id], + 'timestamp': int(time.time()), + 'models': [model_id], + }, + assistant_msg_id: { + 'id': assistant_msg_id, + 'parentId': user_msg_id, + 'role': 'assistant', + 'content': '', + 'done': False, + 'model': model_id, + 'childrenIds': [], + 'timestamp': int(time.time()), + }, + }, + }, + 'messages': [ + {'role': 'user', 'content': prompt}, + ], + 'meta': {'automation_id': automation.id}, + } + ), + ) + + if not chat: + await _record_run(automation.id, 'error', error='Failed to create chat') + return + + # Notify frontend to refresh chat list + from open_webui.socket.main import sio + + await sio.emit( + 'events', + { + 'chat_id': chat.id, + 'message_id': user_msg_id, + 'data': {'type': 'chat:list'}, + }, + room=f'user:{automation.user_id}', + ) + + # Resolve model defaults (frontend does this, backend doesn't) + tool_ids = _resolve_model_tool_ids(app, model_id) + features = _resolve_model_features(app, model_id) + filter_ids = _resolve_model_filter_ids(app, model_id) + + # Resolve terminal from model config + terminal_id = _resolve_model_terminal_id(app, model_id) + + # Build the same payload the frontend sends to /api/chat/completions + form_data = { + 'model': model_id, + 'messages': [{'role': 'user', 'content': prompt}], + 'stream': True, + 'chat_id': chat.id, + 'id': assistant_msg_id, + 'parent_id': user_msg_id, + 'session_id': f'automation:{automation.id}', + 'background_tasks': {}, + } + if tool_ids: + form_data['tool_ids'] = tool_ids + if features: + form_data['features'] = features + if filter_ids: + form_data['filter_ids'] = filter_ids + if terminal_id: + form_data['terminal_id'] = terminal_id + + # Call the full chat completion pipeline (same as POST /api/chat/completions). + # The handler reference is stored on app.state to avoid circular imports. + request = _build_request(app) + await app.state.CHAT_COMPLETION_HANDLER(request, form_data, user=user) + + # Notify user + from open_webui.socket.main import sio + + await sio.emit( + 'automation:result', + { + 'automation_id': automation.id, + 'name': automation.name, + 'chat_id': chat.id, + 'status': 'success', + }, + room=f'user:{automation.user_id}', + ) + + await _record_run(automation.id, 'success', chat_id=chat.id) + + except Exception as e: + log.exception(f'Automation {automation.id} failed') + await _record_run(automation.id, 'error', error=str(e)[:4000]) + + +#################### +# Internals +#################### + + +async def _record_run( + automation_id: str, + status: str, + chat_id: str = None, + error: str = None, +): + """Insert a run record into automation_run.""" + 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 5ce6fffec6..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 @@ -316,6 +316,10 @@ async def chat_completed(request: Request, form_data: dict, user: Any): models = request.app.state.MODELS data = form_data + + if not data.get('id'): + raise Exception('Missing message id') + model_id = data['model'] if model_id not in models: raise Exception('Model not found') @@ -327,6 +331,9 @@ async def chat_completed(request: Request, form_data: dict, user: Any): except Exception as e: raise Exception(f'Error: {e}') + if not data.get('id'): + raise Exception('Missing message id') + metadata = { 'chat_id': data['chat_id'], 'message_id': data['id'], @@ -336,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, @@ -345,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'![{match.group(1)}]({url})' - return match.group(0) +async def convert_markdown_base64_images(request, content: str, metadata, user): + MIN_REPLACEMENT_URL_LENGTH = 1024 + result_parts = [] + last_end = 0 - return MARKDOWN_IMAGE_URL_PATTERN.sub(replace, content) + for match in MARKDOWN_IMAGE_URL_PATTERN.finditer(content): + result_parts.append(content[last_end : match.start()]) + base64_string = match.group(2) + if len(base64_string) > MIN_REPLACEMENT_URL_LENGTH: + url = await get_image_url_from_base64(request, base64_string, metadata, user) + if url: + result_parts.append(f'![{match.group(1)}]({url})') + else: + result_parts.append(match.group(0)) + else: + result_parts.append(match.group(0)) + last_end = match.end() + + result_parts.append(content[last_end:]) + return ''.join(result_parts) def load_b64_audio_data(b64_str): @@ -110,7 +119,7 @@ def load_b64_audio_data(b64_str): return None, None -def upload_audio(request, audio_data, content_type, metadata, user): +async def upload_audio(request, audio_data, content_type, metadata, user): audio_format = mimetypes.guess_extension(content_type) file = UploadFile( file=io.BytesIO(audio_data), @@ -119,7 +128,7 @@ def upload_audio(request, audio_data, content_type, metadata, user): 'content-type': content_type, }, ) - file_item = upload_file_handler( + file_item = await upload_file_handler( request, file=file, metadata=metadata, @@ -130,13 +139,13 @@ def upload_audio(request, audio_data, content_type, metadata, user): return url -def get_audio_url_from_base64(request, base64_audio_string, metadata, user): +async def get_audio_url_from_base64(request, base64_audio_string, metadata, user): if 'data:audio/wav;base64' in base64_audio_string: audio_url = '' # Extract base64 audio data from the line audio_data, content_type = load_b64_audio_data(base64_audio_string) if audio_data is not None: - audio_url = upload_audio( + audio_url = await upload_audio( request, audio_data, content_type, @@ -147,16 +156,16 @@ def get_audio_url_from_base64(request, base64_audio_string, metadata, user): return None -def get_file_url_from_base64(request, base64_file_string, metadata, user): +async def get_file_url_from_base64(request, base64_file_string, metadata, user): if BASE64_IMAGE_URL_PREFIX.match(base64_file_string): - return get_image_url_from_base64(request, base64_file_string, metadata, user) + return await get_image_url_from_base64(request, base64_file_string, metadata, user) elif 'data:audio/wav;base64' in base64_file_string: - return get_audio_url_from_base64(request, base64_file_string, metadata, user) + return await get_audio_url_from_base64(request, base64_file_string, metadata, user) return None -def get_image_base64_from_file_id(id: str) -> Optional[str]: - file = Files.get_file_by_id(id) +async def get_image_base64_from_file_id(id: str) -> Optional[str]: + file = await Files.get_file_by_id(id) if not file: return None diff --git a/backend/open_webui/utils/filter.py b/backend/open_webui/utils/filter.py index df07dea4a1..50b1583088 100644 --- a/backend/open_webui/utils/filter.py +++ b/backend/open_webui/utils/filter.py @@ -10,44 +10,53 @@ from open_webui.models.functions import Functions log = logging.getLogger(__name__) -def get_function_module(request, function_id, load_from_db=True): +async def get_function_module(request, function_id, load_from_db=True): """ Get the function module by its ID. """ - function_module, _, _ = get_function_module_from_cache(request, function_id, load_from_db) + function_module, _, _ = await get_function_module_from_cache(request, function_id, load_from_db) return function_module -def get_sorted_filter_ids(request, model: dict, enabled_filter_ids: list = None): - def get_priority(function_id): +async def get_sorted_filter_ids(request, model: dict, enabled_filter_ids: list = None): + async def get_priority(function_id): try: - function_module = get_function_module(request, function_id) + function_module = await get_function_module(request, function_id) if function_module and hasattr(function_module, 'Valves'): - valves_db = Functions.get_function_valves_by_id(function_id) + valves_db = await Functions.get_function_valves_by_id(function_id) valves = function_module.Valves(**(valves_db if valves_db else {})) return getattr(valves, 'priority', 0) except Exception: pass return 0 - filter_ids = [function.id for function in Functions.get_global_filter_functions()] + filter_ids = [function.id for function in await Functions.get_global_filter_functions()] if 'info' in model and 'meta' in model['info']: filter_ids.extend(model['info']['meta'].get('filterIds', [])) filter_ids = list(set(filter_ids)) - active_filter_ids = {function.id for function in Functions.get_functions_by_type('filter', active_only=True)} + active_filter_ids = {function.id for function in await Functions.get_functions_by_type('filter', active_only=True)} - def get_active_status(filter_id): - function_module = get_function_module(request, filter_id) + async def get_active_status(filter_id): + function_module = await get_function_module(request, filter_id) if getattr(function_module, 'toggle', None): return filter_id in (enabled_filter_ids or set()) return True - active_filter_ids = {filter_id for filter_id in active_filter_ids if get_active_status(filter_id)} + # Pre-compute active status for each filter (async functions can't be used in set comprehensions) + resolved_active = {} + for filter_id in active_filter_ids: + resolved_active[filter_id] = await get_active_status(filter_id) + active_filter_ids = {fid for fid, is_active in resolved_active.items() if is_active} filter_ids = [fid for fid in filter_ids if fid in active_filter_ids] - filter_ids.sort(key=lambda fid: (get_priority(fid), fid)) + + # Pre-compute priorities (async functions can't be used in sort keys) + priorities = {} + for fid in filter_ids: + priorities[fid] = await get_priority(fid) + filter_ids.sort(key=lambda fid: (priorities.get(fid, 0), fid)) return filter_ids @@ -63,7 +72,7 @@ async def process_filter_functions(request, filter_functions, filter_type, form_ if not filter: continue - function_module = get_function_module(request, filter_id, load_from_db=(filter_type != 'stream')) + function_module = await get_function_module(request, filter_id, load_from_db=(filter_type != 'stream')) # Prepare handler function handler = getattr(function_module, filter_type, None) if not handler: @@ -75,7 +84,7 @@ async def process_filter_functions(request, filter_functions, filter_type, form_ # Apply valves to the function if hasattr(function_module, 'valves') and hasattr(function_module, 'Valves'): - valves = Functions.get_function_valves_by_id(filter_id) + valves = await Functions.get_function_valves_by_id(filter_id) function_module.valves = function_module.Valves(**(valves if valves else {})) try: @@ -100,7 +109,7 @@ async def process_filter_functions(request, filter_functions, filter_type, form_ if hasattr(function_module, 'UserValves'): try: params['__user__']['valves'] = function_module.UserValves( - **Functions.get_user_valves_by_id_and_user_id(filter_id, params['__user__']['id']) + **await Functions.get_user_valves_by_id_and_user_id(filter_id, params['__user__']['id']) ) except Exception as e: log.exception(f'Failed to get user values: {e}') diff --git a/backend/open_webui/utils/groups.py b/backend/open_webui/utils/groups.py index 90c4593cec..50099b2ee7 100644 --- a/backend/open_webui/utils/groups.py +++ b/backend/open_webui/utils/groups.py @@ -4,7 +4,7 @@ from open_webui.models.groups import Groups log = logging.getLogger(__name__) -def apply_default_group_assignment( +async def apply_default_group_assignment( default_group_id: str, user_id: str, db=None, @@ -18,6 +18,6 @@ def apply_default_group_assignment( """ if default_group_id: try: - Groups.add_users_to_group(default_group_id, [user_id], db=db) + await Groups.add_users_to_group(default_group_id, [user_id], db=db) except Exception as e: log.error(f'Failed to add user {user_id} to default group {default_group_id}: {e}') diff --git a/backend/open_webui/utils/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 79221f4463..effe4b1637 100644 --- a/backend/open_webui/utils/mcp/client.py +++ b/backend/open_webui/utils/mcp/client.py @@ -1,7 +1,10 @@ import asyncio +import logging from typing import Optional from contextlib import AsyncExitStack +log = logging.getLogger(__name__) + import anyio from mcp import ClientSession @@ -41,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) @@ -56,7 +59,9 @@ class MCPClient: self._streams_context = streamablehttp_client( url, headers=headers, - httpx_client_factory=create_httpx_client if AIOHTTP_CLIENT_SESSION_TOOL_SERVER_SSL else create_insecure_httpx_client, + httpx_client_factory=create_httpx_client + if AIOHTTP_CLIENT_SESSION_TOOL_SERVER_SSL + else create_insecure_httpx_client, ) transport = await exit_stack.enter_async_context(self._streams_context) @@ -134,8 +139,39 @@ class MCPClient: return result_dict async def disconnect(self): - # Clean up and close the session - await self.exit_stack.aclose() + """Clean up and close the session. + + This method is idempotent — calling it multiple times or on a + client that was never connected is safe. It shields the close + operation from CancelledError and adds a timeout so a hung MCP + server cannot block the event loop indefinitely. + """ + exit_stack = self.exit_stack + if exit_stack is None: + return + + # Prevent double-close from concurrent callers + self.exit_stack = None + self.session = None + + try: + await asyncio.wait_for( + asyncio.shield(exit_stack.aclose()), + timeout=5.0, + ) + except asyncio.TimeoutError: + log.warning('MCPClient.disconnect() timed out after 5 s') + except RuntimeError as exc: + # The MCP SDK's streamable_http transport uses anyio task + # groups and async generators internally. When we close + # a session that was interrupted mid-flight these can + # raise RuntimeError ("aclose(): asynchronous generator is + # already running" or "Attempted to exit cancel scope in a + # different task"). Swallowing the error here prevents the + # orphaned coroutines from spinning at 100 % CPU. + log.debug('MCPClient.disconnect() suppressed RuntimeError: %s', exc) + except Exception as exc: + log.debug('MCPClient.disconnect() error: %s', exc) async def __aenter__(self): await self.exit_stack.__aenter__() diff --git a/backend/open_webui/utils/middleware.py b/backend/open_webui/utils/middleware.py index 317a471467..1f708ead25 100644 --- a/backend/open_webui/utils/middleware.py +++ b/backend/open_webui/utils/middleware.py @@ -890,10 +890,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'{body}\n' ) return context_string @@ -937,7 +941,7 @@ def apply_source_context_to_messages( ) -def process_tool_result( +async def process_tool_result( request, tool_function_name, tool_result, @@ -1076,7 +1080,7 @@ def process_tool_result( pass tool_response.append(text) elif item.get('type') in ['image', 'audio']: - file_url = get_file_url_from_base64( + file_url = await get_file_url_from_base64( request, f'data:{item.get("mimeType")};base64,{item.get("data", item.get("blob", ""))}', { @@ -1305,7 +1309,7 @@ async def chat_completion_tools_handler( except Exception as e: tool_result = str(e) - tool_result, tool_result_files, tool_result_embeds = process_tool_result( + tool_result, tool_result_files, tool_result_embeds = await process_tool_result( request, tool_function_name, tool_result, @@ -1603,7 +1607,7 @@ def get_images_from_messages(message_list): return images -def get_image_urls(delta_images, request, metadata, user) -> list[str]: +async def get_image_urls(delta_images, request, metadata, user) -> list[str]: if not isinstance(delta_images, list): return [] @@ -1617,21 +1621,21 @@ def get_image_urls(delta_images, request, metadata, user) -> list[str]: continue if url.startswith('data:image/png;base64'): - url = get_image_url_from_base64(request, url, metadata, user) + url = await get_image_url_from_base64(request, url, metadata, user) image_urls.append(url) return image_urls -def add_file_context(messages: list, chat_id: str, user) -> list: +async def add_file_context(messages: list, chat_id: str, user) -> list: """ Add file URLs to messages for native function calling. """ if not chat_id or chat_id.startswith('local:'): return messages - chat = Chats.get_chat_by_id_and_user_id(chat_id, user.id) + chat = await Chats.get_chat_by_id_and_user_id(chat_id, user.id) if not chat: return messages @@ -1687,7 +1691,7 @@ async def chat_image_generation_handler(request: Request, form_data: dict, extra if chat_id.startswith('local:'): message_list = form_data.get('messages', []) else: - chat = Chats.get_chat_by_id_and_user_id(chat_id, user.id) + chat = await Chats.get_chat_by_id_and_user_id(chat_id, user.id) await __event_emitter__( { 'type': 'status', @@ -2067,12 +2071,12 @@ async def convert_url_images_to_base64(form_data): return form_data -def load_messages_from_db(chat_id: str, message_id: str) -> Optional[list[dict]]: +async def load_messages_from_db(chat_id: str, message_id: str) -> Optional[list[dict]]: """ Load the message chain from DB up to message_id, keeping only LLM-relevant fields (role, content, output). """ - messages_map = Chats.get_messages_map_by_chat_id(chat_id) + messages_map = await Chats.get_messages_map_by_chat_id(chat_id) if not messages_map: return None @@ -2156,7 +2160,7 @@ async def process_chat_payload(request, form_data, user, metadata, model): ) if chat_id and context_message_id and not chat_id.startswith('local:'): - db_messages = load_messages_from_db(chat_id, context_message_id) + db_messages = await load_messages_from_db(chat_id, context_message_id) if db_messages: system_message = get_system_message(form_data.get('messages', [])) form_data['messages'] = [system_message, *db_messages] if system_message else db_messages @@ -2199,8 +2203,8 @@ async def process_chat_payload(request, form_data, user, metadata, model): form_data = await convert_url_images_to_base64(form_data) - event_emitter = get_event_emitter(metadata) - event_caller = get_event_call(metadata) + event_emitter = await get_event_emitter(metadata) + event_caller = await get_event_call(metadata) extra_params = { '__event_emitter__': event_emitter, @@ -2238,14 +2242,14 @@ async def process_chat_payload(request, form_data, user, metadata, model): chat_id = metadata.get('chat_id', None) folder_id = None if chat_id and user: - folder_id = Chats.get_chat_folder_id(chat_id, user.id) + folder_id = await Chats.get_chat_folder_id(chat_id, user.id) # Fallback: use folder_id from metadata (temporary chats have no DB record) if not folder_id: folder_id = metadata.get('folder_id', None) if folder_id and user: - folder = Folders.get_folder_by_id_and_user_id(folder_id, user.id) + folder = await Folders.get_folder_by_id_and_user_id(folder_id, user.id) if folder and folder.data: if 'system_prompt' in folder.data: @@ -2312,8 +2316,8 @@ async def process_chat_payload(request, form_data, user, metadata, model): raise e try: - filter_ids = get_sorted_filter_ids(request, model, metadata.get('filter_ids', [])) - filter_functions = Functions.get_functions_by_ids(filter_ids) + filter_ids = await get_sorted_filter_ids(request, model, metadata.get('filter_ids', [])) + filter_functions = await Functions.get_functions_by_ids(filter_ids) form_data, flags = await process_filter_functions( request=request, @@ -2377,11 +2381,16 @@ async def process_chat_payload(request, form_data, user, metadata, model): ) else: # Native FC: tool docstring can't be dynamic, so inject - # filesystem context into messages for pyodide engine + # filesystem context into the system message for pyodide + # engine. Appending to the system prompt (instead of the + # user message) keeps it in the stable cached prefix so + # providers with prefix caching don't re-bill the full + # conversation on every turn. if engine != 'jupyter': - form_data['messages'] = add_or_update_user_message( + form_data['messages'] = add_or_update_system_message( CODE_INTERPRETER_PYODIDE_PROMPT, form_data['messages'], + append=True, ) tool_ids = form_data.pop('tool_ids', None) @@ -2401,12 +2410,13 @@ async def process_chat_payload(request, form_data, user, metadata, model): if all_skill_ids: from open_webui.models.skills import Skills as SkillsModel - accessible_skill_ids = {s.id for s in SkillsModel.get_skills_by_user_id(user.id, 'read')} - available_skills = [ - s - for sid in all_skill_ids - if sid in accessible_skill_ids and (s := SkillsModel.get_skill_by_id(sid)) and s.is_active - ] + accessible_skill_ids = {s.id for s in await SkillsModel.get_skills_by_user_id(user.id, 'read')} + available_skills = [] + for sid in all_skill_ids: + if sid in accessible_skill_ids: + s = await SkillsModel.get_skill_by_id(sid) + if s and s.is_active: + available_skills.append(s) skill_descriptions = '' for skill in available_skills: @@ -2419,7 +2429,7 @@ async def process_chat_payload(request, form_data, user, metadata, model): ) else: # Model-attached: name+description only - skill_descriptions += f'\n{skill.name}\n{skill.description or ""}\n\n' + skill_descriptions += f'\n{skill.id}\n{skill.name}\n{skill.description or ""}\n\n' if skill_descriptions: form_data['messages'] = add_or_update_system_message( @@ -2443,7 +2453,7 @@ async def process_chat_payload(request, form_data, user, metadata, model): # Get folder files folder_id = file_item.get('id', None) if folder_id: - folder = Folders.get_folder_by_id_and_user_id(folder_id, user.id) + folder = await Folders.get_folder_by_id_and_user_id(folder_id, user.id) if folder and folder.data and 'files' in folder.data: files = [f for f in files if f.get('id', None) != folder_id] files = [*files, *folder.data['files']] @@ -2497,7 +2507,7 @@ async def process_chat_payload(request, form_data, user, metadata, model): continue # Check access control for MCP server - if not has_connection_access(user, mcp_server_connection): + if not await has_connection_access(user, mcp_server_connection): log.warning(f'Access denied to MCP server {server_id} for user {user.id}') continue @@ -2514,7 +2524,7 @@ async def process_chat_payload(request, form_data, user, metadata, model): oauth_token = extra_params.get('__oauth_token__', None) if oauth_token: headers['Authorization'] = f'Bearer {oauth_token.get("access_token", "")}' - elif auth_type == 'oauth_2.1': + elif auth_type in ('oauth_2.1', 'oauth_2.1_static'): try: splits = server_id.split(':') server_id = splits[-1] if len(splits) > 1 else server_id @@ -2558,7 +2568,7 @@ async def process_chat_payload(request, form_data, user, metadata, model): tool_specs = await mcp_clients[server_id].list_tool_specs() for tool_spec in tool_specs: - def make_tool_function(client, function_name): + async def make_tool_function(client, function_name): async def tool_function(**kwargs): return await client.call_tool( function_name, @@ -2572,7 +2582,7 @@ async def process_chat_payload(request, form_data, user, metadata, model): # Skip this function continue - tool_function = make_tool_function(mcp_clients[server_id], tool_spec['name']) + tool_function = await make_tool_function(mcp_clients[server_id], tool_spec['name']) mcp_tools_dict[f'{server_id}_{tool_spec["name"]}'] = { 'spec': { @@ -2666,8 +2676,8 @@ async def process_chat_payload(request, form_data, user, metadata, model): if metadata.get('params', {}).get('function_calling') == 'native' and builtin_tools_enabled: # Add file context to user messages chat_id = metadata.get('chat_id') - form_data['messages'] = add_file_context(form_data.get('messages', []), chat_id, user) - builtin_tools = get_builtin_tools( + form_data['messages'] = await add_file_context(form_data.get('messages', []), chat_id, user) + builtin_tools = await get_builtin_tools( request, { **extra_params, @@ -2682,9 +2692,12 @@ async def process_chat_payload(request, form_data, user, metadata, model): tools_dict[name] = tool_dict if tools_dict: + # Always store resolved tools in metadata so downstream consumers + # (e.g. pipe functions) can access all tools including MCP and builtins. + metadata['tools'] = tools_dict + if metadata.get('params', {}).get('function_calling') == 'native': # If the function calling is native, then call the tools function calling handler - metadata['tools'] = tools_dict form_data['tools'] = [ {'type': 'function', 'function': tool.get('spec', {})} for tool in tools_dict.values() ] @@ -2754,24 +2767,26 @@ async def process_chat_payload(request, form_data, user, metadata, model): return form_data, metadata, events -def get_event_emitter_and_caller(metadata): +async def get_event_emitter_and_caller(metadata): event_emitter = None event_caller = None - if ( - 'session_id' in metadata - and metadata['session_id'] - and 'chat_id' in metadata - and metadata['chat_id'] - and 'message_id' in metadata - and metadata['message_id'] - ): - event_emitter = get_event_emitter(metadata) - event_caller = get_event_call(metadata) + + # event_emitter only needs user_id + chat_id + message_id. + # It broadcasts to user:{user_id} room AND persists to DB, + # so it works for backend-initiated calls (automations, API). + if metadata.get('chat_id') and metadata.get('message_id'): + event_emitter = await get_event_emitter(metadata) + + # event_caller needs session_id — it calls back to a specific + # websocket session (used by direct tools, pyodide code interpreter). + if metadata.get('session_id') and metadata.get('chat_id') and metadata.get('message_id'): + event_caller = await get_event_call(metadata) + return event_emitter, event_caller -def build_chat_response_context(request, form_data, user, model, metadata, tasks, events): - event_emitter, event_caller = get_event_emitter_and_caller(metadata) +async def build_chat_response_context(request, form_data, user, model, metadata, tasks, events): + event_emitter, event_caller = await get_event_emitter_and_caller(metadata) return { 'request': request, 'form_data': form_data, @@ -2859,7 +2874,7 @@ async def background_tasks_handler(ctx): messages = [] if 'chat_id' in metadata and not metadata['chat_id'].startswith('local:'): - messages_map = Chats.get_messages_map_by_chat_id(metadata['chat_id']) + messages_map = await Chats.get_messages_map_by_chat_id(metadata['chat_id']) message = messages_map.get(metadata['message_id']) if messages_map else None message_list = get_message_list(messages_map, metadata['message_id']) @@ -2939,7 +2954,7 @@ async def background_tasks_handler(ctx): ) if not metadata.get('chat_id', '').startswith('local:'): - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], { @@ -2992,7 +3007,7 @@ async def background_tasks_handler(ctx): if not title: title = messages[0].get('content', user_message) - Chats.update_chat_title_by_id(metadata['chat_id'], title) + await Chats.update_chat_title_by_id(metadata['chat_id'], title) await event_emitter( { @@ -3004,7 +3019,7 @@ async def background_tasks_handler(ctx): if title == None and len(messages) == 2 and (not messages_map or len(messages_map) <= 2): title = messages[0].get('content', user_message) - Chats.update_chat_title_by_id(metadata['chat_id'], title) + await Chats.update_chat_title_by_id(metadata['chat_id'], title) await event_emitter( { @@ -3038,7 +3053,7 @@ async def background_tasks_handler(ctx): try: tags = json.loads(tags_string).get('tags', []) - Chats.update_chat_tags_by_id(metadata['chat_id'], tags, user) + await Chats.update_chat_tags_by_id(metadata['chat_id'], tags, user) await event_emitter( { @@ -3073,7 +3088,9 @@ async def non_streaming_chat_response_handler(response, ctx): else: error = str(error) - Chats.upsert_message_to_chat_by_id_and_message_id( + log.error('Provider returned error (non-streaming): %s', error) + + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], { @@ -3089,7 +3106,7 @@ async def non_streaming_chat_response_handler(response, ctx): ) if 'selected_model_id' in response_data: - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], { @@ -3109,7 +3126,7 @@ async def non_streaming_chat_response_handler(response, ctx): } ) - title = Chats.get_chat_title_by_id(metadata['chat_id']) + title = await Chats.get_chat_title_by_id(metadata['chat_id']) # Use output from backend if provided (OR-compliant backends), # otherwise generate from response content @@ -3140,7 +3157,7 @@ async def non_streaming_chat_response_handler(response, ctx): # Save message in the database usage = normalize_usage(response_data.get('usage', {}) or {}) - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], { @@ -3153,8 +3170,8 @@ async def non_streaming_chat_response_handler(response, ctx): ) # Send a webhook notification if the user is not active - if request.app.state.config.ENABLE_USER_WEBHOOKS and not Users.is_user_active(user.id): - webhook_url = Users.get_user_webhook_url_by_id(user.id) + if request.app.state.config.ENABLE_USER_WEBHOOKS and not await Users.is_user_active(user.id): + webhook_url = await Users.get_user_webhook_url_by_id(user.id) if webhook_url: await post_webhook( request.app.state.WEBUI_NAME, @@ -3208,12 +3225,14 @@ async def streaming_chat_response_handler(response, ctx): } filter_functions = [ - Functions.get_function_by_id(filter_id) - for filter_id in get_sorted_filter_ids(request, model, metadata.get('filter_ids', [])) + await Functions.get_function_by_id(filter_id) + for filter_id in await get_sorted_filter_ids(request, model, metadata.get('filter_ids', [])) ] # Standard streaming response handler - if event_emitter and event_caller: + # event_caller is optional — only needed for direct (client-side) tools + # and pyodide code interpreter. Server-side tools work without it. + if event_emitter: task_id = str(uuid4()) # Create a unique task ID. model_id = form_data.get('model', '') @@ -3442,7 +3461,7 @@ async def streaming_chat_response_handler(response, ctx): return output, end_flag - message = Chats.get_message_by_id_and_message_id(metadata['chat_id'], metadata['message_id']) + message = await Chats.get_message_by_id_and_message_id(metadata['chat_id'], metadata['message_id']) tool_calls = [] @@ -3504,7 +3523,7 @@ async def streaming_chat_response_handler(response, ctx): ) # Save message in the database - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], { @@ -3574,7 +3593,7 @@ async def streaming_chat_response_handler(response, ctx): if 'selected_model_id' in data: model_id = data['selected_model_id'] - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], { @@ -3639,6 +3658,17 @@ async def streaming_chat_response_handler(response, ctx): if not choices: error = data.get('error', {}) if error: + log.error('Provider returned error (streaming): %s', error) + try: + await Chats.upsert_message_to_chat_by_id_and_message_id( + metadata['chat_id'], + metadata['message_id'], + { + 'error': {'content': error}, + }, + ) + except Exception: + pass await event_emitter( { 'type': 'chat:completion', @@ -3747,10 +3777,10 @@ async def streaming_chat_response_handler(response, ctx): } ) - image_urls = get_image_urls(delta.get('images', []), request, metadata, user) + image_urls = await get_image_urls(delta.get('images', []), request, metadata, user) if image_urls: image_file_list = [{'type': 'image', 'url': url} for url in image_urls] - message_files = Chats.add_message_files_by_id_and_message_id( + message_files = await Chats.add_message_files_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], image_file_list, @@ -3832,7 +3862,7 @@ async def streaming_chat_response_handler(response, ctx): ) if ENABLE_CHAT_RESPONSE_BASE64_IMAGE_URL_CONVERSION: - value = convert_markdown_base64_images( + value = await convert_markdown_base64_images( request, value, { @@ -3948,7 +3978,7 @@ async def streaming_chat_response_handler(response, ctx): if ENABLE_REALTIME_CHAT_SAVE: # Save message in the database - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], { @@ -4044,11 +4074,12 @@ async def streaming_chat_response_handler(response, ctx): if responses_api_tool_calls: tool_calls.append(_split_tool_calls(responses_api_tool_calls)) + try: + await stream_body_handler(response, form_data) + finally: if response.background: await response.background() - await stream_body_handler(response, form_data) - tool_call_retries = 0 tool_call_sources = [] # Track citation sources from tool results all_tool_call_sources = [] # Accumulated sources across all iterations @@ -4169,7 +4200,7 @@ async def streaming_chat_response_handler(response, ctx): ) else: - tool_function = get_updated_tool_function( + tool_function = await get_updated_tool_function( function=tool['callable'], extra_params={ '__messages__': form_data.get('messages', []), @@ -4182,7 +4213,7 @@ async def streaming_chat_response_handler(response, ctx): except Exception as e: tool_result = str(e) - tool_result, tool_result_files, tool_result_embeds = process_tool_result( + tool_result, tool_result_files, tool_result_embeds = await process_tool_result( request, tool_function_name, tool_result, @@ -4472,7 +4503,7 @@ async def streaming_chat_response_handler(response, ctx): BLOCKED_MODULES = {CODE_INTERPRETER_BLOCKED_MODULES} _real_import = builtins.__import__ - def restricted_import(name, globals=None, locals=None, fromlist=(), level=0): + async def restricted_import(name, globals=None, locals=None, fromlist=(), level=0): if name.split('.')[0] in BLOCKED_MODULES: importer_name = globals.get('__name__') if globals else None if importer_name == '__main__': @@ -4526,7 +4557,7 @@ async def streaming_chat_response_handler(response, ctx): stdoutLines = stdout.split('\n') for idx, line in enumerate(stdoutLines): if re.match(r'data:image/\w+;base64', line): - image_url = get_image_url_from_base64( + image_url = await get_image_url_from_base64( request, line, metadata, @@ -4543,7 +4574,7 @@ async def streaming_chat_response_handler(response, ctx): resultLines = result.split('\n') for idx, line in enumerate(resultLines): if re.match(r'data:image/\w+;base64', line): - image_url = get_image_url_from_base64( + image_url = await get_image_url_from_base64( request, line, metadata, @@ -4608,17 +4639,18 @@ async def streaming_chat_response_handler(response, ctx): if item.get('status') == 'in_progress': item['status'] = 'completed' - title = Chats.get_chat_title_by_id(metadata['chat_id']) + title = await Chats.get_chat_title_by_id(metadata['chat_id']) data = { 'done': True, 'content': serialize_output(output), 'output': output, 'title': title, + **({'usage': usage} if usage else {}), } if not ENABLE_REALTIME_CHAT_SAVE: # Save message in the database - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], { @@ -4629,21 +4661,21 @@ async def streaming_chat_response_handler(response, ctx): }, ) elif usage: - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], {'done': True, 'usage': usage}, ) else: - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], {'done': True}, ) # Send a webhook notification if the user is not active - if request.app.state.config.ENABLE_USER_WEBHOOKS and not Users.is_user_active(user.id): - webhook_url = Users.get_user_webhook_url_by_id(user.id) + if request.app.state.config.ENABLE_USER_WEBHOOKS and not await Users.is_user_active(user.id): + webhook_url = await Users.get_user_webhook_url_by_id(user.id) if webhook_url: await post_webhook( request.app.state.WEBUI_NAME, @@ -4671,7 +4703,7 @@ async def streaming_chat_response_handler(response, ctx): if not ENABLE_REALTIME_CHAT_SAVE: # Save message in the database - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], { @@ -4681,7 +4713,7 @@ async def streaming_chat_response_handler(response, ctx): }, ) else: - Chats.upsert_message_to_chat_by_id_and_message_id( + await Chats.upsert_message_to_chat_by_id_and_message_id( metadata['chat_id'], metadata['message_id'], {'done': True}, diff --git a/backend/open_webui/utils/misc.py b/backend/open_webui/utils/misc.py index 68a52643d7..984ebde6a7 100644 --- a/backend/open_webui/utils/misc.py +++ b/backend/open_webui/utils/misc.py @@ -950,9 +950,11 @@ async def cleanup_response( session: Optional[aiohttp.ClientSession], ): if response: - response.close() + if not response.closed: + await response.close() if session: - await session.close() + if not session.closed: + await session.close() async def stream_wrapper(response, session, content_handler=None): @@ -1006,18 +1008,15 @@ def stream_chunks_handler(stream: aiohttp.StreamReader): skip_mode = False yield line else: - yield b'data: {}' - yield b'\n' + yield b'data: {}\n' else: # Normal mode: check if line exceeds limit if len(line) > max_buffer_size: skip_mode = True - yield b'data: {}' - yield b'\n' + yield b'data: {}\n' log.info(f'Skip mode triggered, line size: {len(line)}') else: - yield line - yield b'\n' + yield line + b'\n' # Save the last incomplete fragment buffer = lines[-1] @@ -1031,7 +1030,6 @@ def stream_chunks_handler(stream: aiohttp.StreamReader): # Process remaining buffer data if buffer and not skip_mode: - yield buffer - yield b'\n' + yield buffer + b'\n' return yield_safe_stream_chunks() diff --git a/backend/open_webui/utils/models.py b/backend/open_webui/utils/models.py index b57c74744c..c8ebc190e5 100644 --- a/backend/open_webui/utils/models.py +++ b/backend/open_webui/utils/models.py @@ -130,13 +130,13 @@ async def get_all_models(request, refresh: bool = False, user: UserModel = None) ] models = models + arena_models - global_action_ids = {function.id for function in Functions.get_global_action_functions()} - enabled_action_ids = {function.id for function in Functions.get_functions_by_type('action', active_only=True)} + global_action_ids = {function.id for function in await Functions.get_global_action_functions()} + enabled_action_ids = {function.id for function in await Functions.get_functions_by_type('action', active_only=True)} - global_filter_ids = {function.id for function in Functions.get_global_filter_functions()} - enabled_filter_ids = {function.id for function in Functions.get_functions_by_type('filter', active_only=True)} + global_filter_ids = {function.id for function in await Functions.get_global_filter_functions()} + enabled_filter_ids = {function.id for function in await Functions.get_functions_by_type('filter', active_only=True)} - custom_models = Models.get_all_models() + custom_models = await Models.get_all_models() # Single O(1) lookup: Ollama base names first, then exact IDs (exact wins). base_model_lookup = {} @@ -278,14 +278,14 @@ async def get_all_models(request, refresh: bool = False, user: UserModel = None) all_function_ids.update(global_action_ids) all_function_ids.update(global_filter_ids) - functions_by_id = {f.id: f for f in Functions.get_functions_by_ids(list(all_function_ids))} + functions_by_id = {f.id: f for f in await Functions.get_functions_by_ids(list(all_function_ids))} # Pre-warm the function module cache once per unique function ID. # This ensures each function's DB freshness check runs exactly once, # not once per (model × function) pair. for function_id in all_function_ids: try: - get_function_module_from_cache(request, function_id) + await get_function_module_from_cache(request, function_id) except Exception as e: log.info(f'Failed to load function module for {function_id}: {e}') @@ -312,7 +312,7 @@ async def get_all_models(request, refresh: bool = False, user: UserModel = None) # Batch-fetch all function valves in one query to avoid N+1 DB hits # inside get_action_priority (previously called per action × per model). - all_function_valves = Functions.get_function_valves_by_ids(list(all_function_ids)) + all_function_valves = await Functions.get_function_valves_by_ids(list(all_function_ids)) def get_action_priority(action_id): try: @@ -377,11 +377,11 @@ async def get_all_models(request, refresh: bool = False, user: UserModel = None) return models -def check_model_access(user, model, db=None): +async def check_model_access(user, model, db=None): if model.get('arena'): meta = model.get('info', {}).get('meta', {}) access_grants = meta.get('access_grants', []) - if not has_access( + if not await has_access( user.id, permission='read', access_grants=access_grants, @@ -389,12 +389,12 @@ def check_model_access(user, model, db=None): ): raise Exception('Model not found') else: - model_info = Models.get_model_by_id(model.get('id'), db=db) + model_info = await Models.get_model_by_id(model.get('id'), db=db) if not model_info: raise Exception('Model not found') elif not ( user.id == model_info.user_id - or AccessGrants.has_access( + or await AccessGrants.has_access( user_id=user.id, resource_type='model', resource_id=model_info.id, @@ -405,7 +405,7 @@ def check_model_access(user, model, db=None): raise Exception('Model not found') -def get_filtered_models(models, user, db=None): +async def get_filtered_models(models, user, db=None): # Filter out models that the user does not have access to if ( user.role == 'user' or (user.role == 'admin' and not BYPASS_ADMIN_ACCESS_CONTROL) @@ -418,10 +418,10 @@ def get_filtered_models(models, user, db=None): if info: model_infos[model['id']] = info - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id, db=db)} + user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)} # Batch-fetch accessible resource IDs in a single query instead of N has_access calls - accessible_model_ids = AccessGrants.get_accessible_resource_ids( + accessible_model_ids = await AccessGrants.get_accessible_resource_ids( user_id=user.id, resource_type='model', resource_ids=list(model_infos.keys()), @@ -435,7 +435,7 @@ def get_filtered_models(models, user, db=None): if model.get('arena'): meta = model.get('info', {}).get('meta', {}) access_grants = meta.get('access_grants', []) - if has_access( + if await has_access( user.id, permission='read', access_grants=access_grants, diff --git a/backend/open_webui/utils/oauth.py b/backend/open_webui/utils/oauth.py index e4a97327a2..6020100eb4 100644 --- a/backend/open_webui/utils/oauth.py +++ b/backend/open_webui/utils/oauth.py @@ -18,6 +18,7 @@ from typing import Literal import aiohttp from authlib.integrations.starlette_client import OAuth +from authlib.jose.errors import BadSignatureError from authlib.oidc.core import UserInfo from fastapi import ( HTTPException, @@ -75,11 +76,13 @@ from open_webui.env import ( ENABLE_OAUTH_EMAIL_FALLBACK, OAUTH_CLIENT_INFO_ENCRYPTION_KEY, OAUTH_MAX_SESSIONS_PER_USER, + REDIS_KEY_PREFIX, ) from open_webui.utils.misc import parse_duration from open_webui.utils.auth import get_password_hash, create_token from open_webui.utils.webhook import post_webhook from open_webui.utils.groups import apply_default_group_assignment +from open_webui.retrieval.web.utils import validate_url from mcp.shared.auth import ( OAuthClientMetadata as MCPOAuthClientMetadata, @@ -522,6 +525,7 @@ class OAuthClientManager: 'client_id': oauth_client_info.client_id, 'client_secret': oauth_client_info.client_secret, 'client_kwargs': { + 'follow_redirects': True, **({'scope': oauth_client_info.scope} if oauth_client_info.scope else {}), **( {'token_endpoint_auth_method': oauth_client_info.token_endpoint_auth_method} @@ -698,7 +702,7 @@ class OAuthClientManager: """ try: # Get the OAuth session - session = OAuthSessions.get_session_by_provider_and_user_id(client_id, user_id) + session = await OAuthSessions.get_session_by_provider_and_user_id(client_id, user_id) if not session: log.warning(f'No OAuth session found for user {user_id}, client_id {client_id}') return None @@ -712,7 +716,7 @@ class OAuthClientManager: log.warning( f'Token refresh failed for user {user_id}, client_id {session.provider}, deleting session {session.id}' ) - OAuthSessions.delete_session_by_id(session.id) + await OAuthSessions.delete_session_by_id(session.id) return None return session.token @@ -736,7 +740,7 @@ class OAuthClientManager: if refreshed_token: # Update the session with new token data - session = OAuthSessions.update_session_by_id(session.id, refreshed_token) + session = await OAuthSessions.update_session_by_id(session.id, refreshed_token) log.info(f'Successfully refreshed token for session {session.id}') return session.token else: @@ -882,12 +886,12 @@ class OAuthClientManager: token['expires_at'] = datetime.now().timestamp() + token['expires_in'] # Clean up any existing sessions for this user/client_id first - sessions = OAuthSessions.get_sessions_by_user_id(user_id) + sessions = await OAuthSessions.get_sessions_by_user_id(user_id) for session in sessions: if session.provider == client_id: - OAuthSessions.delete_session_by_id(session.id) + await OAuthSessions.delete_session_by_id(session.id) - session = OAuthSessions.create_session( + session = await OAuthSessions.create_session( user_id=user_id, provider=client_id, token=token, @@ -961,7 +965,7 @@ class OAuthManager: """ try: # Get the OAuth session - session = OAuthSessions.get_session_by_id_and_user_id(session_id, user_id) + session = await OAuthSessions.get_session_by_id_and_user_id(session_id, user_id) if not session: log.warning(f'No OAuth session found for user {user_id}, session {session_id}') return None @@ -975,7 +979,7 @@ class OAuthManager: log.warning( f'Token refresh failed for user {user_id}, provider {session.provider}, deleting session {session.id}' ) - OAuthSessions.delete_session_by_id(session.id) + await OAuthSessions.delete_session_by_id(session.id) return None return session.token @@ -1000,7 +1004,7 @@ class OAuthManager: if refreshed_token: # Update the session with new token data - session = OAuthSessions.update_session_by_id(session.id, refreshed_token) + session = await OAuthSessions.update_session_by_id(session.id, refreshed_token) log.info(f'Successfully refreshed token for session {session.id}') return session.token else: @@ -1100,16 +1104,19 @@ class OAuthManager: log.error(f'Exception during token refresh for provider {provider}: {e}') return None - def get_user_role(self, user, user_data): - user_count = Users.get_num_users() + async def get_user_role(self, user, user_data): + user_count = await Users.get_num_users() if user and user_count == 1: # If the user is the only user, assign the role "admin" - actually repairs role for single user on login log.debug('Assigning the only user the admin role') return 'admin' if not user and user_count == 0: - # If there are no users, assign the role "admin", as the first user will be an admin - log.debug('Assigning the first user the admin role') - return 'admin' + # First-user bootstrap: skip role management gating so the + # instance can be initialized. We intentionally return the + # default role here (not 'admin') — admin promotion happens + # race-safely *after* insert via get_num_users() == 1. + log.debug('First user bootstrap: using default role (admin promotion deferred to post-insert)') + return auth_manager_config.DEFAULT_USER_ROLE if auth_manager_config.ENABLE_OAUTH_ROLE_MANAGEMENT: log.debug('Running OAUTH Role management') @@ -1183,7 +1190,7 @@ class OAuthManager: return role - def update_user_groups(self, user, user_data, default_permissions, db=None): + async def update_user_groups(self, user, user_data, default_permissions, db=None): log.debug('Running OAUTH Group management') oauth_claim = auth_manager_config.OAUTH_GROUPS_CLAIM @@ -1212,8 +1219,8 @@ class OAuthManager: else: user_oauth_groups = [] - user_current_groups: list[GroupModel] = Groups.get_groups_by_member_id(user.id, db=db) - all_available_groups: list[GroupModel] = Groups.get_all_groups(db=db) + user_current_groups: list[GroupModel] = await Groups.get_groups_by_member_id(user.id, db=db) + all_available_groups: list[GroupModel] = await Groups.get_all_groups(db=db) # Create groups if they don't exist and creation is enabled if auth_manager_config.ENABLE_OAUTH_GROUP_CREATION: @@ -1221,7 +1228,7 @@ class OAuthManager: all_group_names = {g.name for g in all_available_groups} groups_created = False # Determine creator ID: Prefer admin, fallback to current user if no admin exists - admin_user = Users.get_super_admin_user() + admin_user = await Users.get_super_admin_user() creator_id = admin_user.id if admin_user else user.id log.debug(f'Using creator ID {creator_id} for potential group creation.') @@ -1236,7 +1243,7 @@ class OAuthManager: data={'config': {'share': auth_manager_config.OAUTH_GROUP_DEFAULT_SHARE}}, ) # Use determined creator ID (admin or fallback to current user) - created_group = Groups.insert_new_group(creator_id, new_group_form, db=db) + created_group = await Groups.insert_new_group(creator_id, new_group_form, db=db) if created_group: log.info( f"Successfully created group '{group_name}' with ID {created_group.id} using creator ID {creator_id}" @@ -1251,7 +1258,7 @@ class OAuthManager: # Refresh the list of all available groups if any were created if groups_created: - all_available_groups = Groups.get_all_groups(db=db) + all_available_groups = await Groups.get_all_groups(db=db) log.debug('Refreshed list of all available groups after creation.') log.debug(f'Oauth Groups claim: {oauth_claim}') @@ -1268,14 +1275,14 @@ class OAuthManager: ): # Remove group from user log.debug(f'Removing user from group {group_model.name} as it is no longer in their oauth groups') - Groups.remove_users_from_group(group_model.id, [user.id], db=db) + await Groups.remove_users_from_group(group_model.id, [user.id], db=db) # In case a group is created, but perms are never assigned to the group by hitting "save" group_permissions = group_model.permissions if not group_permissions: group_permissions = default_permissions - Groups.update_group_by_id( + await Groups.update_group_by_id( id=group_model.id, form_data=GroupUpdateForm( name=group_model.name, @@ -1297,14 +1304,14 @@ class OAuthManager: # Add user to group log.debug(f'Adding user to group {group_model.name} as it was found in their oauth groups') - Groups.add_users_to_group(group_model.id, [user.id], db=db) + await Groups.add_users_to_group(group_model.id, [user.id], db=db) # In case a group is created, but perms are never assigned to the group by hitting "save" group_permissions = group_model.permissions if not group_permissions: group_permissions = default_permissions - Groups.update_group_by_id( + await Groups.update_group_by_id( id=group_model.id, form_data=GroupUpdateForm( name=group_model.name, @@ -1329,6 +1336,8 @@ class OAuthManager: return '/user.png' try: + validate_url(picture_url) + get_kwargs = {} if access_token: get_kwargs['headers'] = { @@ -1354,12 +1363,12 @@ class OAuthManager: if provider not in OAUTH_PROVIDERS: raise HTTPException(404) # If the provider has a custom redirect URL, use that, otherwise automatically generate one - redirect_uri = OAUTH_PROVIDERS[provider].get('redirect_uri') or request.url_for( - 'oauth_login_callback', provider=provider - ) client = self.get_client(provider) if client is None: raise HTTPException(404) + redirect_uri = (client.server_metadata or {}).get('redirect_uri') or request.url_for( + 'oauth_login_callback', provider=provider + ) kwargs = {} if auth_manager_config.OAUTH_AUDIENCE: @@ -1385,6 +1394,27 @@ class OAuthManager: try: token = await client.authorize_access_token(request, **auth_params) + except BadSignatureError: + # The IdP likely rotated its signing keys and the cached JWKS + # is stale. Evict the cached key set so the next attempt + # fetches fresh keys from the jwks_uri. + log.warning( + 'OIDC bad_signature for provider %s — evicting cached JWKS and retrying', + provider, + ) + if hasattr(client, 'server_metadata') and isinstance(client.server_metadata, dict): + client.server_metadata.pop('jwks', None) + try: + token = await client.authorize_access_token(request, **auth_params) + except Exception as retry_exc: + detailed_error = _build_oauth_callback_error_message(retry_exc) + log.warning( + 'OAuth callback error during authorize_access_token retry for provider %s: %s', + provider, + detailed_error, + exc_info=True, + ) + raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED) except Exception as e: detailed_error = _build_oauth_callback_error_message(e) log.warning( @@ -1483,20 +1513,20 @@ class OAuthManager: raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED) # Check if the user exists - user = Users.get_user_by_oauth_sub(provider, sub, db=db) + user = await Users.get_user_by_oauth_sub(provider, sub, db=db) if not user: # If the user does not exist, check if merging is enabled if auth_manager_config.OAUTH_MERGE_ACCOUNTS_BY_EMAIL: # Check if the user exists by email - user = Users.get_user_by_email(email, db=db) + user = await Users.get_user_by_email(email, db=db) if user: # Update the user with the new oauth sub - Users.update_user_oauth_by_id(user.id, provider, sub, db=db) + await Users.update_user_oauth_by_id(user.id, provider, sub, db=db) if user: - determined_role = self.get_user_role(user, user_data) + determined_role = await self.get_user_role(user, user_data) if user.role != determined_role: - Users.update_user_role_by_id(user.id, determined_role, db=db) + await Users.update_user_role_by_id(user.id, determined_role, db=db) # Update the user object in memory as well, # to avoid problems with the ENABLE_OAUTH_GROUP_MANAGEMENT check below user.role = determined_role @@ -1506,7 +1536,7 @@ class OAuthManager: if username_claim: new_name = user_data.get(username_claim) if new_name and new_name != user.name: - Users.update_user_by_id(user.id, {'name': new_name}, db=db) + await Users.update_user_by_id(user.id, {'name': new_name}, db=db) user.name = new_name log.debug(f'Updated name for user {user.email}') @@ -1515,13 +1545,13 @@ class OAuthManager: if email_claim: new_email = user_data.get(email_claim) if new_email and new_email.lower() != user.email.lower(): - existing_user = Users.get_user_by_email(new_email, db=db) + existing_user = await Users.get_user_by_email(new_email, db=db) if existing_user: log.error( f'Cannot update email to {new_email} for user {user.id} because it is already taken.' ) else: - Auths.update_email_by_id(user.id, new_email.lower(), db=db) + await Auths.update_email_by_id(user.id, new_email.lower(), db=db) user.email = new_email.lower() log.debug(f'Updated email for user {user.id}') @@ -1537,13 +1567,13 @@ class OAuthManager: new_picture_url, token.get('access_token') ) if processed_picture_url != user.profile_image_url: - Users.update_user_profile_image_url_by_id(user.id, processed_picture_url, db=db) + await Users.update_user_profile_image_url_by_id(user.id, processed_picture_url, db=db) log.debug(f'Updated profile picture for user {user.email}') else: # If the user does not exist, check if signups are enabled if auth_manager_config.ENABLE_OAUTH_SIGNUP: # Check if an existing user with the same email already exists - existing_user = Users.get_user_by_email(email, db=db) + existing_user = await Users.get_user_by_email(email, db=db) if existing_user: raise HTTPException(400, detail=ERROR_MESSAGES.EMAIL_TAKEN) @@ -1563,16 +1593,26 @@ class OAuthManager: log.warning('Username claim is missing, using email as name') name = email - user = Auths.insert_new_auth( + user = await Auths.insert_new_auth( email=email, password=get_password_hash(str(uuid.uuid4())), # Random password, not used name=name, profile_image_url=picture_url, - role=self.get_user_role(None, user_data), + role=await self.get_user_role(None, user_data), oauth=oauth_data, db=db, ) + if not user: + raise HTTPException(500, detail=ERROR_MESSAGES.CREATE_USER_ERROR) + + # Atomically check if this is the only user *after* the + # insert to avoid TOCTOU race on first-user registration. + # Matches signup_handler pattern. + if await Users.get_num_users(db=db) == 1: + await Users.update_user_role_by_id(user.id, 'admin', db=db) + user = await Users.get_user_by_id(user.id, db=db) + if auth_manager_config.WEBHOOK_URL: await post_webhook( WEBUI_NAME, @@ -1585,7 +1625,7 @@ class OAuthManager: }, ) - apply_default_group_assignment(request.app.state.config.DEFAULT_GROUP_ID, user.id, db=db) + await apply_default_group_assignment(request.app.state.config.DEFAULT_GROUP_ID, user.id, db=db) else: raise HTTPException( @@ -1598,7 +1638,7 @@ class OAuthManager: expires_delta=parse_duration(auth_manager_config.JWT_EXPIRES_IN), ) if auth_manager_config.ENABLE_OAUTH_GROUP_MANAGEMENT: - self.update_user_groups( + await self.update_user_groups( user=user, user_data=user_data, default_permissions=request.app.state.config.USER_PERMISSIONS, @@ -1658,7 +1698,7 @@ class OAuthManager: # Enforce max concurrent sessions per user/provider to prevent # unbounded growth while allowing multi-device usage - sessions = OAuthSessions.get_sessions_by_user_id(user.id, db=db) + sessions = await OAuthSessions.get_sessions_by_user_id(user.id, db=db) provider_sessions = sorted( [session for session in sessions if session.provider == provider], key=lambda session: session.created_at, @@ -1667,9 +1707,9 @@ class OAuthManager: # Keep the newest sessions up to the limit, prune the rest if len(provider_sessions) >= OAUTH_MAX_SESSIONS_PER_USER: for old_session in provider_sessions[OAUTH_MAX_SESSIONS_PER_USER - 1 :]: - OAuthSessions.delete_session_by_id(old_session.id, db=db) + await OAuthSessions.delete_session_by_id(old_session.id, db=db) - session = OAuthSessions.create_session( + session = await OAuthSessions.create_session( user_id=user.id, provider=provider, token=token, @@ -1683,7 +1723,7 @@ class OAuthManager: httponly=True, samesite=WEBUI_AUTH_COOKIE_SAME_SITE, secure=WEBUI_AUTH_COOKIE_SECURE, - **({'max_age': cookie_max_age, 'expires': cookie_expires} if cookie_max_age is not None else {}), + **({'max_age': cookie_max_age} if cookie_max_age is not None else {}), ) log.info(f'Stored OAuth session server-side for user {user.id}, provider {provider}') @@ -1693,3 +1733,186 @@ class OAuthManager: log.error(f'Failed to store OAuth session server-side: {e}') return response + + async def handle_backchannel_logout(self, request, db=None): + """ + Handle an OIDC Back-Channel Logout request. + Validates the logout_token, identifies the user, revokes their + sessions via Redis, and deletes their OAuth sessions. + Returns a JSONResponse per the OIDC Back-Channel Logout 1.0 spec. + """ + import jwt as pyjwt + from fastapi.responses import JSONResponse + + # 1. Extract logout_token from form body + try: + form = await request.form() + logout_token = form.get('logout_token') + except Exception: + logout_token = None + + if not logout_token: + return JSONResponse( + status_code=400, + content={'error': 'invalid_request', 'error_description': 'Missing logout_token parameter'}, + ) + + # 2. Peek at unverified issuer to match against configured providers + try: + unverified_claims = pyjwt.decode(logout_token, options={'verify_signature': False}) + token_issuer = unverified_claims.get('iss') + except Exception as e: + log.warning(f'Back-channel logout: cannot decode logout_token: {e}') + return JSONResponse( + status_code=400, + content={'error': 'invalid_request', 'error_description': 'Malformed logout_token'}, + ) + + if not token_issuer: + return JSONResponse( + status_code=400, + content={'error': 'invalid_request', 'error_description': 'logout_token missing iss claim'}, + ) + + # 3. Find the configured provider whose issuer matches the token + matched_provider = None + matched_client_id = None + matched_jwks_uri = None + matched_issuer = None + + for provider_name in OAUTH_PROVIDERS: + server_metadata_url = self.get_server_metadata_url(provider_name) + if not server_metadata_url: + continue + + try: + async with aiohttp.ClientSession(trust_env=True) as session: + async with session.get(server_metadata_url, ssl=AIOHTTP_CLIENT_SESSION_SSL) as r: + if r.status != 200: + continue + oidc_config = await r.json() + + provider_issuer = oidc_config.get('issuer') + if provider_issuer and provider_issuer == token_issuer: + client = self.get_client(provider_name) + matched_provider = provider_name + matched_client_id = client.client_id if client else None + matched_jwks_uri = oidc_config.get('jwks_uri') + matched_issuer = provider_issuer + break + except Exception as e: + log.debug(f'Back-channel logout: error checking provider {provider_name}: {e}') + continue + + if not matched_provider or not matched_client_id or not matched_jwks_uri: + log.warning(f'Back-channel logout: no configured provider matches issuer {token_issuer}') + return JSONResponse( + status_code=400, + content={ + 'error': 'invalid_request', + 'error_description': 'No configured provider matches token issuer', + }, + ) + + # 4. Validate the logout_token signature and claims + try: + jwks_client = pyjwt.PyJWKClient(matched_jwks_uri) + signing_key = jwks_client.get_signing_key_from_jwt(logout_token) + + claims = pyjwt.decode( + logout_token, + signing_key.key, + algorithms=['RS256', 'RS384', 'RS512', 'ES256', 'ES384', 'ES512'], + audience=matched_client_id, + issuer=matched_issuer, + options={ + 'require': ['iss', 'aud', 'iat', 'events'], + }, + ) + except pyjwt.InvalidTokenError as e: + log.warning(f'Back-channel logout: invalid logout_token: {e}') + return JSONResponse( + status_code=400, + content={'error': 'invalid_request', 'error_description': f'Invalid logout_token: {e}'}, + ) + except Exception as e: + log.error(f'Back-channel logout: error validating logout_token: {e}') + return JSONResponse( + status_code=400, + content={'error': 'invalid_request', 'error_description': 'Failed to validate logout_token'}, + ) + + # 5. Validate events claim per spec + events = claims.get('events', {}) + if 'http://schemas.openid.net/event/backchannel-logout' not in events: + log.warning('Back-channel logout: missing required backchannel-logout event claim') + return JSONResponse( + status_code=400, + content={'error': 'invalid_request', 'error_description': 'Missing backchannel-logout event claim'}, + ) + + # 6. Per spec, back-channel logout tokens MUST NOT contain a nonce + if 'nonce' in claims: + log.warning('Back-channel logout: logout_token contains nonce (rejected per spec)') + return JSONResponse( + status_code=400, + content={'error': 'invalid_request', 'error_description': 'logout_token must not contain nonce'}, + ) + + # 7. Extract sub and/or sid — at least one must be present + sub = claims.get('sub') + sid = claims.get('sid') + + if not sub and not sid: + log.warning('Back-channel logout: logout_token contains neither sub nor sid') + return JSONResponse( + status_code=400, + content={'error': 'invalid_request', 'error_description': 'logout_token must contain sub or sid'}, + ) + + # 8. Identify users to log out + users_to_logout = [] + if sub: + user = await Users.get_user_by_oauth_sub(matched_provider, sub, db=db) + if user: + users_to_logout.append(user) + + if not users_to_logout and sid: + log.info(f'Back-channel logout: no user found by sub, sid-based lookup not yet supported (sid={sid})') + + if not users_to_logout: + log.info(f'Back-channel logout: no matching user for provider={matched_provider}, sub={sub}, sid={sid}') + return JSONResponse(status_code=200, content={}) + + # 9. Revoke tokens and delete sessions + redis = request.app.state.redis + if not redis: + log.warning( + 'Back-channel logout: Redis not configured, cannot revoke JWT tokens. ' + 'OAuth sessions will be deleted but existing JWTs will remain valid until expiry.' + ) + + revoked_count = 0 + for user in users_to_logout: + sessions = await OAuthSessions.get_sessions_by_user_id(user.id, db=db) + for oauth_session in sessions: + await OAuthSessions.delete_session_by_id(oauth_session.id, db=db) + + if redis: + revocation_key = f'{REDIS_KEY_PREFIX}:auth:user:{user.id}:revoked_at' + await redis.set( + revocation_key, + str(int(time.time())), + ex=60 * 60 * 24 * 30, + ) + revoked_count += 1 + + log.info( + f'Back-channel logout: revoked sessions for user {user.id} ' + f'(email={user.email}, provider={matched_provider}, sessions_deleted={len(sessions)})' + ) + + log.info( + f'Back-channel logout: completed for {len(users_to_logout)} user(s), {revoked_count} revocation(s) set' + ) + return JSONResponse(status_code=200, content={}) diff --git a/backend/open_webui/utils/plugin.py b/backend/open_webui/utils/plugin.py index 46622e21ae..84671bbd3b 100644 --- a/backend/open_webui/utils/plugin.py +++ b/backend/open_webui/utils/plugin.py @@ -199,16 +199,16 @@ def replace_imports(content): # May the intent of the one who wrote it survive every # import and transformation, as a deed survives the generations. -def load_tool_module_by_id(tool_id, content=None): +async def load_tool_module_by_id(tool_id, content=None): if content is None: - tool = Tools.get_tool_by_id(tool_id) + tool = await Tools.get_tool_by_id(tool_id) if not tool: raise Exception(f'Toolkit not found: {tool_id}') content = tool.content content = replace_imports(content) - Tools.update_tool_by_id(tool_id, {'content': content}) + await Tools.update_tool_by_id(tool_id, {'content': content}) else: frontmatter = extract_frontmatter(content) # Install required packages found within the frontmatter @@ -245,15 +245,15 @@ def load_tool_module_by_id(tool_id, content=None): os.unlink(temp_file.name) -def load_function_module_by_id(function_id: str, content: str | None = None): +async def load_function_module_by_id(function_id: str, content: str | None = None): if content is None: - function = Functions.get_function_by_id(function_id) + function = await Functions.get_function_by_id(function_id) if not function: raise Exception(f'Function not found: {function_id}') content = function.content content = replace_imports(content) - Functions.update_function_by_id(function_id, {'content': content}) + await Functions.update_function_by_id(function_id, {'content': content}) else: frontmatter = extract_frontmatter(content) install_frontmatter_requirements(frontmatter.get('requirements', '')) @@ -290,16 +290,16 @@ def load_function_module_by_id(function_id: str, content: str | None = None): # Cleanup by removing the module in case of error del sys.modules[module_name] - Functions.update_function_by_id(function_id, {'is_active': False}) + await Functions.update_function_by_id(function_id, {'is_active': False}) raise e finally: os.unlink(temp_file.name) -def get_tool_module_from_cache(request, tool_id, load_from_db=True): +async def get_tool_module_from_cache(request, tool_id, load_from_db=True): if load_from_db: # Always load from the database by default - tool = Tools.get_tool_by_id(tool_id) + tool = await Tools.get_tool_by_id(tool_id) if not tool: raise Exception(f'Tool not found: {tool_id}') content = tool.content @@ -308,7 +308,7 @@ def get_tool_module_from_cache(request, tool_id, load_from_db=True): if new_content != content: content = new_content # Update the tool content in the database - Tools.update_tool_by_id(tool_id, {'content': content}) + await Tools.update_tool_by_id(tool_id, {'content': content}) if (hasattr(request.app.state, 'TOOL_CONTENTS') and tool_id in request.app.state.TOOL_CONTENTS) and ( hasattr(request.app.state, 'TOOLS') and tool_id in request.app.state.TOOLS @@ -316,12 +316,12 @@ def get_tool_module_from_cache(request, tool_id, load_from_db=True): if request.app.state.TOOL_CONTENTS[tool_id] == content: return request.app.state.TOOLS[tool_id], None - tool_module, frontmatter = load_tool_module_by_id(tool_id, content) + tool_module, frontmatter = await load_tool_module_by_id(tool_id, content) else: if hasattr(request.app.state, 'TOOLS') and tool_id in request.app.state.TOOLS: return request.app.state.TOOLS[tool_id], None - tool_module, frontmatter = load_tool_module_by_id(tool_id) + tool_module, frontmatter = await load_tool_module_by_id(tool_id) if not hasattr(request.app.state, 'TOOLS'): request.app.state.TOOLS = {} @@ -335,13 +335,13 @@ def get_tool_module_from_cache(request, tool_id, load_from_db=True): return tool_module, frontmatter -def get_function_module_from_cache(request, function_id, load_from_db=True): +async def get_function_module_from_cache(request, function_id, load_from_db=True): if load_from_db: # Always load from the database by default # This is useful for hooks like "inlet" or "outlet" where the content might change # and we want to ensure the latest content is used. - function = Functions.get_function_by_id(function_id) + function = await Functions.get_function_by_id(function_id) if not function: raise Exception(f'Function not found: {function_id}') content = function.content @@ -350,7 +350,7 @@ def get_function_module_from_cache(request, function_id, load_from_db=True): if new_content != content: content = new_content # Update the function content in the database - Functions.update_function_by_id(function_id, {'content': content}) + await Functions.update_function_by_id(function_id, {'content': content}) if ( hasattr(request.app.state, 'FUNCTION_CONTENTS') and function_id in request.app.state.FUNCTION_CONTENTS @@ -358,7 +358,7 @@ def get_function_module_from_cache(request, function_id, load_from_db=True): if request.app.state.FUNCTION_CONTENTS[function_id] == content: return request.app.state.FUNCTIONS[function_id], None, None - function_module, function_type, frontmatter = load_function_module_by_id(function_id, content) + function_module, function_type, frontmatter = await load_function_module_by_id(function_id, content) else: # Load from cache (e.g. "stream" hook) # This is useful for performance reasons @@ -366,7 +366,7 @@ def get_function_module_from_cache(request, function_id, load_from_db=True): if hasattr(request.app.state, 'FUNCTIONS') and function_id in request.app.state.FUNCTIONS: return request.app.state.FUNCTIONS[function_id], None, None - function_module, function_type, frontmatter = load_function_module_by_id(function_id) + function_module, function_type, frontmatter = await load_function_module_by_id(function_id) if not hasattr(request.app.state, 'FUNCTIONS'): request.app.state.FUNCTIONS = {} @@ -404,7 +404,7 @@ def install_frontmatter_requirements(requirements: str): log.info('No requirements found in frontmatter.') -def install_tool_and_function_dependencies(): +async def install_tool_and_function_dependencies(): """ Install all dependencies for all admin tools and active functions. @@ -412,8 +412,8 @@ def install_tool_and_function_dependencies(): and then installing them using pip. Duplicates or similar version specifications are handled by pip as much as possible. """ - function_list = Functions.get_functions(active_only=True) - tool_list = Tools.get_tools() + function_list = await Functions.get_functions(active_only=True) + tool_list = await Tools.get_tools() all_dependencies = '' try: diff --git a/backend/open_webui/utils/redis.py b/backend/open_webui/utils/redis.py index 55d08147a9..cb570cb45a 100644 --- a/backend/open_webui/utils/redis.py +++ b/backend/open_webui/utils/redis.py @@ -9,7 +9,9 @@ import redis from open_webui.env import ( REDIS_CLUSTER, + REDIS_HEALTH_CHECK_INTERVAL, REDIS_SOCKET_CONNECT_TIMEOUT, + REDIS_SOCKET_KEEPALIVE, REDIS_SENTINEL_HOSTS, REDIS_SENTINEL_MAX_RETRY_COUNT, REDIS_SENTINEL_PORT, @@ -36,7 +38,7 @@ class SentinelRedisProxy: def _master(self): return self._sentinel.master_for(self._service, **self._kw) - def __getattr__(self, item): + async def __getattr__(self, item): master = self._master() orig_attr = getattr(master, item) @@ -191,6 +193,14 @@ def get_redis_connection( connection = None + connect_timeout_kwargs = ( + {'socket_connect_timeout': REDIS_SOCKET_CONNECT_TIMEOUT} if REDIS_SOCKET_CONNECT_TIMEOUT is not None else {} + ) + + keepalive_kwargs = {'socket_keepalive': True} if REDIS_SOCKET_KEEPALIVE else {} + + health_check_kwargs = {'health_check_interval': REDIS_HEALTH_CHECK_INTERVAL} if REDIS_HEALTH_CHECK_INTERVAL else {} + if async_mode: import redis.asyncio as redis @@ -205,6 +215,8 @@ def get_redis_connection( password=redis_config['password'], decode_responses=decode_responses, socket_connect_timeout=REDIS_SOCKET_CONNECT_TIMEOUT, + **keepalive_kwargs, + **health_check_kwargs, ) connection = SentinelRedisProxy( sentinel, @@ -214,9 +226,21 @@ def get_redis_connection( elif redis_cluster: if not redis_url: raise ValueError('Redis URL must be provided for cluster mode.') - return redis.cluster.RedisCluster.from_url(redis_url, decode_responses=decode_responses) + return redis.cluster.RedisCluster.from_url( + redis_url, + decode_responses=decode_responses, + **connect_timeout_kwargs, + **keepalive_kwargs, + **health_check_kwargs, + ) elif redis_url: - connection = redis.from_url(redis_url, decode_responses=decode_responses) + connection = redis.from_url( + redis_url, + decode_responses=decode_responses, + **connect_timeout_kwargs, + **keepalive_kwargs, + **health_check_kwargs, + ) else: import redis @@ -230,6 +254,8 @@ def get_redis_connection( password=redis_config['password'], decode_responses=decode_responses, socket_connect_timeout=REDIS_SOCKET_CONNECT_TIMEOUT, + **keepalive_kwargs, + **health_check_kwargs, ) connection = SentinelRedisProxy( sentinel, @@ -239,9 +265,21 @@ def get_redis_connection( elif redis_cluster: if not redis_url: raise ValueError('Redis URL must be provided for cluster mode.') - return redis.cluster.RedisCluster.from_url(redis_url, decode_responses=decode_responses) + return redis.cluster.RedisCluster.from_url( + redis_url, + decode_responses=decode_responses, + **connect_timeout_kwargs, + **keepalive_kwargs, + **health_check_kwargs, + ) elif redis_url: - connection = redis.Redis.from_url(redis_url, decode_responses=decode_responses) + connection = redis.Redis.from_url( + redis_url, + decode_responses=decode_responses, + **connect_timeout_kwargs, + **keepalive_kwargs, + **health_check_kwargs, + ) _CONNECTION_CACHE[cache_key] = connection return connection diff --git a/backend/open_webui/utils/session_pool.py b/backend/open_webui/utils/session_pool.py new file mode 100644 index 0000000000..e91579f4af --- /dev/null +++ b/backend/open_webui/utils/session_pool.py @@ -0,0 +1,113 @@ +"""Shared aiohttp ClientSession pool. + +Instead of creating a new ClientSession (and TCPConnector) per request, +callers acquire a long-lived session from this module. The pool manages +a single TCPConnector with configurable limits, enabling TCP/SSL connection +reuse, shared DNS cache, and bounded concurrency. + +All pool parameters are configurable via environment variables: + - AIOHTTP_POOL_CONNECTIONS (default 100) — max total connections + - AIOHTTP_POOL_CONNECTIONS_PER_HOST (default 30) — per-host limit + - AIOHTTP_POOL_DNS_TTL (default 300) — DNS cache TTL in seconds + +Usage: + from open_webui.utils.session_pool import get_session, cleanup_response + + session = await get_session() + r = await session.request(...) + # When done with the *response* (not the session): + await cleanup_response(r) + +IMPORTANT: Callers must NOT close the shared session. Only the response +needs cleanup. The session is closed once during application shutdown +via ``close_session()``. +""" + +import logging +from typing import Optional + +import aiohttp + +from open_webui.env import ( + AIOHTTP_CLIENT_TIMEOUT, + AIOHTTP_POOL_CONNECTIONS, + AIOHTTP_POOL_CONNECTIONS_PER_HOST, + AIOHTTP_POOL_DNS_TTL, +) + +log = logging.getLogger(__name__) + +_session: Optional[aiohttp.ClientSession] = None + + +async def get_session() -> aiohttp.ClientSession: + """Return the shared aiohttp ClientSession, creating it lazily.""" + global _session + if _session is None or _session.closed: + connector_kwargs = { + 'ttl_dns_cache': AIOHTTP_POOL_DNS_TTL, + 'enable_cleanup_closed': True, + } + if AIOHTTP_POOL_CONNECTIONS is not None: + connector_kwargs['limit'] = AIOHTTP_POOL_CONNECTIONS + else: + connector_kwargs['limit'] = 0 # aiohttp: 0 = unlimited + if AIOHTTP_POOL_CONNECTIONS_PER_HOST is not None: + connector_kwargs['limit_per_host'] = AIOHTTP_POOL_CONNECTIONS_PER_HOST + else: + connector_kwargs['limit_per_host'] = 0 # aiohttp: 0 = unlimited + connector = aiohttp.TCPConnector(**connector_kwargs) + timeout = aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT) + _session = aiohttp.ClientSession( + connector=connector, + timeout=timeout, + trust_env=True, + ) + log.info( + 'Created shared aiohttp session pool (limit=%s, per_host=%s, dns_ttl=%d)', + AIOHTTP_POOL_CONNECTIONS or 'unlimited', + AIOHTTP_POOL_CONNECTIONS_PER_HOST or 'unlimited', + AIOHTTP_POOL_DNS_TTL, + ) + return _session + + +async def close_session(): + """Close the shared session. Called during application shutdown.""" + global _session + if _session and not _session.closed: + await _session.close() + log.info('Closed shared aiohttp session pool') + _session = None + + +async def cleanup_response( + response: Optional[aiohttp.ClientResponse], + session: Optional[aiohttp.ClientSession] = None, +): + """Release and close an aiohttp response, optionally closing the session. + + When using the shared pool, ``session`` should be ``None`` (the pool + session is never closed per-request). When a caller creates its own + one-off session, pass it here to close it after the response. + """ + if response: + if not response.closed: + await response.close() + if session: + if not session.closed: + await session.close() + + +async def stream_wrapper(response, session=None, content_handler=None): + """Wrap a stream to ensure cleanup happens even if streaming is interrupted. + + This is more reliable than BackgroundTask which may not run if the client + disconnects. When using the shared pool, ``session`` should be ``None``. + """ + try: + stream = content_handler(response.content) if content_handler else response.content + async for chunk in stream: + yield chunk + finally: + await cleanup_response(response, session) diff --git a/backend/open_webui/utils/task.py b/backend/open_webui/utils/task.py index 203c429d22..15213b1f03 100644 --- a/backend/open_webui/utils/task.py +++ b/backend/open_webui/utils/task.py @@ -19,7 +19,7 @@ def get_task_model_id(default_model_id: str, task_model: str, task_model_externa # Set the task model task_model_id = default_model_id # Check if the user has a custom task model and use that model - if models[task_model_id].get('connection_type') == 'local': + if models.get(task_model_id, {}).get('connection_type') == 'local': if task_model and task_model in models: task_model_id = task_model else: diff --git a/backend/open_webui/utils/telemetry/instrumentors.py b/backend/open_webui/utils/telemetry/instrumentors.py index 394e7178d6..fe8e9ba799 100644 --- a/backend/open_webui/utils/telemetry/instrumentors.py +++ b/backend/open_webui/utils/telemetry/instrumentors.py @@ -7,8 +7,8 @@ from aiohttp import ( TraceRequestEndParams, TraceRequestExceptionParams, ) -from chromadb.telemetry.opentelemetry.fastapi import instrument_fastapi from fastapi import FastAPI +from opentelemetry.instrumentation.fastapi import FastAPIInstrumentor from opentelemetry.instrumentation.httpx import ( HTTPXClientInstrumentor, RequestInfo, @@ -176,7 +176,7 @@ class Instrumentor(BaseInstrumentor): return [] def _instrument(self, **kwargs): - instrument_fastapi(app=self.app) + FastAPIInstrumentor.instrument_app(app=self.app) SQLAlchemyInstrumentor().instrument(engine=self.db_engine) RedisInstrumentor().instrument(request_hook=redis_request_hook) RequestsInstrumentor().instrument(request_hook=requests_hook, response_hook=response_hook) diff --git a/backend/open_webui/utils/telemetry/metrics.py b/backend/open_webui/utils/telemetry/metrics.py index 4c43de3342..26216b6ca4 100644 --- a/backend/open_webui/utils/telemetry/metrics.py +++ b/backend/open_webui/utils/telemetry/metrics.py @@ -124,16 +124,16 @@ def setup_metrics(app: FastAPI, resource: Resource) -> None: unit='ms', ) - def observe_active_users( + async def observe_active_users( options: metrics.CallbackOptions, ) -> Sequence[metrics.Observation]: return [ metrics.Observation( - value=Users.get_active_user_count(), + value=await Users.get_active_user_count(), ) ] - def observe_total_registered_users( + async def observe_total_registered_users( options: metrics.CallbackOptions, ) -> Sequence[metrics.Observation]: # IMPORTANT: Use get_num_users() for efficient COUNT(*) query. @@ -141,7 +141,7 @@ def setup_metrics(app: FastAPI, resource: Resource) -> None: # causing connection pool exhaustion on high-latency databases (e.g., Aurora). return [ metrics.Observation( - value=Users.get_num_users() or 0, + value=await Users.get_num_users() or 0, ) ] @@ -159,10 +159,10 @@ def setup_metrics(app: FastAPI, resource: Resource) -> None: callbacks=[observe_active_users], ) - def observe_users_active_today( + async def observe_users_active_today( options: metrics.CallbackOptions, ) -> Sequence[metrics.Observation]: - return [metrics.Observation(value=Users.get_num_users_active_today())] + return [metrics.Observation(value=await Users.get_num_users_active_today())] meter.create_observable_gauge( name='webui.users.active.today', diff --git a/backend/open_webui/utils/tools.py b/backend/open_webui/utils/tools.py index 377a81d749..da817c4741 100644 --- a/backend/open_webui/utils/tools.py +++ b/backend/open_webui/utils/tools.py @@ -51,6 +51,7 @@ from open_webui.env import ( ENABLE_FORWARD_USER_INFO_HEADERS, FORWARD_SESSION_INFO_HEADER_CHAT_ID, FORWARD_SESSION_INFO_HEADER_MESSAGE_ID, + REDIS_KEY_PREFIX, ) from open_webui.utils.headers import include_user_info_headers from open_webui.tools.builtin import ( @@ -85,17 +86,26 @@ from open_webui.tools.builtin import ( view_file, view_knowledge_file, view_skill, - tasks, + create_tasks, + update_task, + create_automation, + update_automation, + list_automations, + toggle_automation, + delete_automation, ) import copy +from open_webui.utils.access_control import has_permission log = logging.getLogger(__name__) # Let no function be called without need, and let what # it yields justify the cost of running it. -def get_async_tool_function_and_apply_extra_params(function: Callable, extra_params: dict) -> Callable[..., Awaitable]: +async def get_async_tool_function_and_apply_extra_params( + function: Callable, extra_params: dict +) -> Callable[..., Awaitable]: sig = inspect.signature(function) extra_params = {k: v for k, v in extra_params.items() if k in sig.parameters} partial_func = partial(function, **extra_params) @@ -132,13 +142,13 @@ def get_async_tool_function_and_apply_extra_params(function: Callable, extra_par return new_function -def get_updated_tool_function(function: Callable, extra_params: dict): +async def get_updated_tool_function(function: Callable, extra_params: dict): # Get the original function and merge updated params __function__ = getattr(function, '__function__', None) __extra_params__ = getattr(function, '__extra_params__', None) if __function__ is not None and __extra_params__ is not None: - return get_async_tool_function_and_apply_extra_params( + return await get_async_tool_function_and_apply_extra_params( __function__, {**__extra_params__, **extra_params}, ) @@ -154,16 +164,16 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr tools_dict = {} # Get user's group memberships for access control checks - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)} + user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)} for tool_id in tool_ids: - tool = Tools.get_tool_by_id(tool_id) + tool = await Tools.get_tool_by_id(tool_id) if tool: # Check access control for local tools if ( not (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL) and tool.user_id != user.id - and not AccessGrants.has_access( + and not await AccessGrants.has_access( user_id=user.id, resource_type='tool', resource_id=tool.id, @@ -176,7 +186,7 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr module = request.app.state.TOOLS.get(tool_id, None) if module is None: - module, _ = load_tool_module_by_id(tool_id) + module, _ = await load_tool_module_by_id(tool_id) request.app.state.TOOLS[tool_id] = module __user__ = { @@ -185,11 +195,11 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr # Set valves for the tool if hasattr(module, 'valves') and hasattr(module, 'Valves'): - valves = Tools.get_tool_valves_by_id(tool_id) or {} + valves = await Tools.get_tool_valves_by_id(tool_id) or {} module.valves = module.Valves(**valves) if hasattr(module, 'UserValves'): __user__['valves'] = module.UserValves( # type: ignore - **Tools.get_user_valves_by_id_and_user_id(tool_id, user.id) + **await Tools.get_user_valves_by_id_and_user_id(tool_id, user.id) ) for spec in tool.specs: @@ -207,7 +217,7 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr # convert to function that takes only model params and inserts custom params function_name = spec['name'] tool_function = getattr(module, function_name) - callable = get_async_tool_function_and_apply_extra_params( + callable = await get_async_tool_function_and_apply_extra_params( tool_function, { **extra_params, @@ -279,7 +289,7 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr tool_server_connection = connections[tool_server_idx] # Check access control for tool server - if not has_connection_access(user, tool_server_connection, user_group_ids): + if not await has_connection_access(user, tool_server_connection, user_group_ids): log.warning(f'Access denied to tool server {server_id} for user {user.id}') continue @@ -333,7 +343,7 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr if metadata and metadata.get('message_id'): headers[FORWARD_SESSION_INFO_HEADER_MESSAGE_ID] = metadata.get('message_id') - def make_tool_function(function_name, tool_server_data, headers): + async def make_tool_function(function_name, tool_server_data, headers): async def tool_function(**kwargs): return await execute_tool_server( url=tool_server_data['url'], @@ -346,9 +356,9 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr return tool_function - tool_function = make_tool_function(function_name, tool_server_data, headers) + tool_function = await make_tool_function(function_name, tool_server_data, headers) - callable = get_async_tool_function_and_apply_extra_params( + callable = await get_async_tool_function_and_apply_extra_params( tool_function, {}, ) @@ -375,7 +385,7 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr return tools_dict -def get_builtin_tools( +async def get_builtin_tools( request: Request, extra_params: dict, features: dict = None, model: dict = None ) -> dict[str, dict]: """ @@ -397,6 +407,18 @@ def get_builtin_tools( builtin_tools = model.get('info', {}).get('meta', {}).get('builtinTools', {}) return builtin_tools.get(category, True) + # Helper to check user-level feature permission (admins always pass) + user = extra_params.get('__user__', {}) + + async def has_user_permission(feature_key: str) -> bool: + if user.get('role') == 'admin': + return True + return await has_permission( + user.get('id', ''), + f'features.{feature_key}', + request.app.state.config.USER_PERMISSIONS, + ) + # Time utilities - available for date calculations if is_builtin_tool_enabled('time'): builtin_functions.extend([get_current_timestamp, calculate_timestamp]) @@ -440,7 +462,11 @@ def get_builtin_tools( builtin_functions.extend([search_chats, view_chat]) # Add memory tools if builtin category enabled AND enabled for this chat - if is_builtin_tool_enabled('memory') and (features.get('memory') or get_model_capability('memory', False)): + if ( + is_builtin_tool_enabled('memory') + and (features.get('memory') or get_model_capability('memory', False)) + and await has_user_permission('memories') + ): builtin_functions.extend( [ search_memories, @@ -457,6 +483,7 @@ def get_builtin_tools( and getattr(request.app.state.config, 'ENABLE_WEB_SEARCH', False) and get_model_capability('web_search') and features.get('web_search') + and await has_user_permission('web_search') ): builtin_functions.extend([search_web, fetch_url]) @@ -466,6 +493,7 @@ def get_builtin_tools( and getattr(request.app.state.config, 'ENABLE_IMAGE_GENERATION', False) and get_model_capability('image_generation') and features.get('image_generation') + and await has_user_permission('image_generation') ): builtin_functions.append(generate_image) if ( @@ -473,6 +501,7 @@ def get_builtin_tools( and getattr(request.app.state.config, 'ENABLE_IMAGE_EDIT', False) and get_model_capability('image_generation') and features.get('image_generation') + and await has_user_permission('image_generation') ): builtin_functions.append(edit_image) @@ -482,15 +511,24 @@ def get_builtin_tools( and getattr(request.app.state.config, 'ENABLE_CODE_INTERPRETER', True) and get_model_capability('code_interpreter') and features.get('code_interpreter') + and await has_user_permission('code_interpreter') ): builtin_functions.append(execute_code) - # Notes tools - search, view, create, and update user's notes (if builtin category enabled AND notes enabled globally) - if is_builtin_tool_enabled('notes') and getattr(request.app.state.config, 'ENABLE_NOTES', False): + # Notes tools - search, view, create, and update user's notes + if ( + is_builtin_tool_enabled('notes') + and getattr(request.app.state.config, 'ENABLE_NOTES', False) + and await has_user_permission('notes') + ): builtin_functions.extend([search_notes, view_note, write_note, replace_note_content]) - # Channels tools - search channels and messages (if builtin category enabled AND channels enabled globally) - if is_builtin_tool_enabled('channels') and getattr(request.app.state.config, 'ENABLE_CHANNELS', False): + # Channels tools - search channels and messages + if ( + is_builtin_tool_enabled('channels') + and getattr(request.app.state.config, 'ENABLE_CHANNELS', False) + and await has_user_permission('channels') + ): builtin_functions.extend( [ search_channels, @@ -506,10 +544,16 @@ def get_builtin_tools( # Task management - break down complex work into trackable steps if is_builtin_tool_enabled('tasks'): - builtin_functions.append(tasks) + builtin_functions.extend([create_tasks, update_task]) + + # Automation tools - create and manage scheduled automations from chat + if is_builtin_tool_enabled('automations') and await has_user_permission('automations'): + builtin_functions.extend( + [create_automation, update_automation, list_automations, toggle_automation, delete_automation] + ) for func in builtin_functions: - callable = get_async_tool_function_and_apply_extra_params( + callable = await get_async_tool_function_and_apply_extra_params( func, { '__request__': request, @@ -696,20 +740,31 @@ def get_tool_specs(tool_module: object) -> list[dict]: return specs -def resolve_schema(schema, components): +def resolve_schema(schema, components, resolved_schemas=None): """ Recursively resolves a JSON schema using OpenAPI components. """ if not schema: return {} + if resolved_schemas is None: + resolved_schemas = set() + if '$ref' in schema: ref_path = schema['$ref'] + schema_name = ref_path.split('/')[-1] + + if schema_name in resolved_schemas: + # Avoid infinite recursion on circular references + return {} + + resolved_schemas.add(schema_name) + ref_parts = ref_path.strip('#/').split('/') resolved = components for part in ref_parts[1:]: # Skip the initial 'components' resolved = resolved.get(part, {}) - return resolve_schema(resolved, components) + return resolve_schema(resolved, components, resolved_schemas) resolved_schema = copy.deepcopy(schema) @@ -757,7 +812,7 @@ def convert_openapi_to_tool_payload(openapi_spec): if not description: description = param.get('description') or '' if param_schema.get('enum') and isinstance(param_schema.get('enum'), list): - description += f'. Possible values: {", ".join(param_schema.get("enum"))}' + description += f'. Possible values: {", ".join(str(v) for v in param_schema.get("enum"))}' param_property = { 'type': param_schema.get('type') or 'string', 'description': description, @@ -800,7 +855,9 @@ async def set_tool_servers(request: Request): request.app.state.TOOL_SERVERS = await get_tool_servers_data(request.app.state.config.TOOL_SERVER_CONNECTIONS) if request.app.state.redis is not None: - await request.app.state.redis.set('tool_servers', json.dumps(request.app.state.TOOL_SERVERS)) + await request.app.state.redis.set( + f'{REDIS_KEY_PREFIX}:tool_servers', json.dumps(request.app.state.TOOL_SERVERS) + ) return request.app.state.TOOL_SERVERS @@ -809,7 +866,7 @@ async def get_tool_servers(request: Request): tool_servers = [] if request.app.state.redis is not None: try: - tool_servers = json.loads(await request.app.state.redis.get('tool_servers')) + tool_servers = json.loads(await request.app.state.redis.get(f'{REDIS_KEY_PREFIX}:tool_servers')) request.app.state.TOOL_SERVERS = tool_servers except Exception as e: log.error(f'Error fetching tool_servers from Redis: {e}') @@ -935,7 +992,9 @@ async def set_terminal_servers(request: Request): ) if request.app.state.redis is not None: - await request.app.state.redis.set('terminal_servers', json.dumps(request.app.state.TERMINAL_SERVERS)) + await request.app.state.redis.set( + f'{REDIS_KEY_PREFIX}:terminal_servers', json.dumps(request.app.state.TERMINAL_SERVERS) + ) return request.app.state.TERMINAL_SERVERS @@ -945,7 +1004,7 @@ async def get_terminal_servers(request: Request): terminal_servers = [] if request.app.state.redis is not None: try: - terminal_servers = json.loads(await request.app.state.redis.get('terminal_servers')) + terminal_servers = json.loads(await request.app.state.redis.get(f'{REDIS_KEY_PREFIX}:terminal_servers')) request.app.state.TERMINAL_SERVERS = terminal_servers except Exception as e: log.error(f'Error fetching terminal_servers from Redis: {e}') @@ -975,8 +1034,8 @@ async def get_terminal_tools( log.warning(f'Terminal server not found: {terminal_id}') return {} - user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)} - if not has_connection_access(user, connection, user_group_ids): + user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)} + if not await has_connection_access(user, connection, user_group_ids): log.warning(f'Access denied to terminal {terminal_id} for user {user.id}') return {} @@ -1028,7 +1087,7 @@ async def get_terminal_tools( tool_spec.get('description', '') + f'\n\nThe current working directory is: {terminal_cwd}' ) - def make_tool_function(fn_name, srv_data, hdrs, cks): + async def make_tool_function(fn_name, srv_data, hdrs, cks): async def tool_function(**kwargs): return await execute_tool_server( url=srv_data['url'], @@ -1041,8 +1100,8 @@ async def get_terminal_tools( return tool_function - tool_function = make_tool_function(function_name, server_data, headers, cookies) - callable = get_async_tool_function_and_apply_extra_params(tool_function, {}) + tool_function = await make_tool_function(function_name, server_data, headers, cookies) + callable = await get_async_tool_function_and_apply_extra_params(tool_function, {}) tools_dict[function_name] = { 'tool_id': f'terminal:{terminal_id}', diff --git a/backend/open_webui/utils/validate.py b/backend/open_webui/utils/validate.py index c2064de257..1e98b41105 100644 --- a/backend/open_webui/utils/validate.py +++ b/backend/open_webui/utils/validate.py @@ -1,9 +1,30 @@ """Validation utilities for user-supplied input.""" -# Known static asset paths used as default profile images -_ALLOWED_STATIC_PATHS = ( - '/user.png', - '/static/favicon.png', +import re +from urllib.parse import urlparse + +# Matches the OWUI-generated profile image route. ``[^/?#]+`` accepts +# any user-ID without allowing path-traversal or query/fragment injection, +# and the ``$`` anchor rejects trailing path components. +_USER_PROFILE_IMAGE_RE = re.compile(r'^/api/v1/users/[^/?#]+/profile/image$') + +# Validates MIME type and structure of base64 data URIs. Only the prefix +# is checked — validating the full base64 payload would mean running a +# regex across megabytes of data on every Pydantic instantiation for zero +# security benefit (corrupt base64 simply renders a broken image, same as +# a 404 URL). SVG is intentionally excluded: it can carry embedded scripts. +_SAFE_DATA_URI_RE = re.compile(r'^data:image/(png|jpeg|gif|webp);base64,', re.IGNORECASE) + +# Exact relative paths accepted as profile images. These are the only +# static-asset paths OWUI itself assigns; no prefix/wildcard matching is +# used so that arbitrary relative paths cannot trigger authenticated GETs +# against internal endpoints when rendered as ```` sources. +_SAFE_STATIC_PATHS = frozenset( + { + '/user.png', + '/favicon.png', + '/static/favicon.png', + } ) @@ -13,24 +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) + - 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 - - 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/app.css b/src/app.css index b7b70aeeb3..93417b680a 100644 --- a/src/app.css +++ b/src/app.css @@ -243,6 +243,19 @@ select { animation: smoothFadeIn 0.2s forwards; } +@keyframes fade-in-token { + from { + opacity: 0; + } + to { + opacity: 1; + } +} + +.fade-in-token { + animation: fade-in-token 100ms ease-out; +} + .katex-mathml { display: none; } diff --git a/src/lib/apis/auths/index.ts b/src/lib/apis/auths/index.ts index 1fd22494b5..c501a36ed7 100644 --- a/src/lib/apis/auths/index.ts +++ b/src/lib/apis/auths/index.ts @@ -413,6 +413,9 @@ export const updateUserProfile = async (token: string, profile: object) => { .catch((err) => { console.error(err); error = err.detail; + if (Array.isArray(error)) { + error = error.map((e: { msg?: string }) => e.msg).join('; '); + } return null; }); diff --git a/src/lib/apis/automations/index.ts b/src/lib/apis/automations/index.ts new file mode 100644 index 0000000000..a79fe1ddc8 --- /dev/null +++ b/src/lib/apis/automations/index.ts @@ -0,0 +1,300 @@ +import { WEBUI_API_BASE_URL } from '$lib/constants'; + +export type AutomationTerminalConfig = { + server_id: string; + cwd?: string; +}; + +export type AutomationData = { + prompt: string; + model_id: string; + rrule: string; + terminal?: AutomationTerminalConfig; +}; + +export type AutomationForm = { + name: string; + data: AutomationData; + meta?: { + system_prompt?: string; + temperature?: number; + max_tokens?: number; + webhook?: string; + }; + is_active?: boolean; +}; + +export type AutomationRunModel = { + id: string; + automation_id: string; + chat_id: string | null; + status: string; + error: string | null; + created_at: number; +}; + +export type AutomationResponse = { + id: string; + user_id: string; + name: string; + data: AutomationData; + meta: Record | null; + is_active: boolean; + last_run_at: number | null; + next_run_at: number | null; + + created_at: number; + updated_at: number; + last_run: AutomationRunModel | null; + next_runs: number[] | null; +}; + +export const getAutomationItems = async ( + token: string, + query: string | null, + status: string | null, + page: number +): Promise<{ items: AutomationResponse[]; total: number }> => { + let error = null; + + const searchParams = new URLSearchParams(); + if (query) { + searchParams.append('query', query); + } + if (status && status !== 'all') { + searchParams.append('status', status); + } + if (page) { + searchParams.append('page', page.toString()); + } + + const res = await fetch(`${WEBUI_API_BASE_URL}/automations/list?${searchParams.toString()}`, { + method: 'GET', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + } + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err.detail; + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const createAutomation = async (token: string, form: AutomationForm) => { + let error = null; + + const res = await fetch(`${WEBUI_API_BASE_URL}/automations/create`, { + method: 'POST', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + }, + body: JSON.stringify(form) + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err.detail; + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const getAutomationById = async (token: string, id: string) => { + let error = null; + + const res = await fetch(`${WEBUI_API_BASE_URL}/automations/${id}`, { + method: 'GET', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + } + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err.detail; + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const updateAutomationById = async (token: string, id: string, form: AutomationForm) => { + let error = null; + + const res = await fetch(`${WEBUI_API_BASE_URL}/automations/${id}/update`, { + method: 'POST', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + }, + body: JSON.stringify(form) + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err.detail; + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const toggleAutomationById = async (token: string, id: string) => { + let error = null; + + const res = await fetch(`${WEBUI_API_BASE_URL}/automations/${id}/toggle`, { + method: 'POST', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + } + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err.detail; + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const runAutomationById = async (token: string, id: string) => { + let error = null; + + const res = await fetch(`${WEBUI_API_BASE_URL}/automations/${id}/run`, { + method: 'POST', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + } + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err.detail; + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const deleteAutomationById = async (token: string, id: string) => { + let error = null; + + const res = await fetch(`${WEBUI_API_BASE_URL}/automations/${id}/delete`, { + method: 'DELETE', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + } + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err.detail; + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const getAutomationRuns = async ( + token: string, + id: string, + skip: number = 0, + limit: number = 50 +) => { + let error = null; + + const res = await fetch( + `${WEBUI_API_BASE_URL}/automations/${id}/runs?skip=${skip}&limit=${limit}`, + { + method: 'GET', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + } + } + ) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err.detail; + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; diff --git a/src/lib/apis/evaluations/index.ts b/src/lib/apis/evaluations/index.ts index a3af6e80bb..bfb6955bcf 100644 --- a/src/lib/apis/evaluations/index.ts +++ b/src/lib/apis/evaluations/index.ts @@ -161,13 +161,48 @@ export const getModelHistory = async (token: string = '', modelId: string, days: return res; }; -export const getFeedbackItems = async (token: string = '', orderBy, direction, page) => { +export const getFeedbackModelIds = async (token: string = '') => { + let error = null; + + const res = await fetch(`${WEBUI_API_BASE_URL}/evaluations/feedbacks/models`, { + method: 'GET', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + } + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err.detail; + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const getFeedbackItems = async ( + token: string = '', + orderBy, + direction, + page, + modelId: string = '' +) => { let error = null; const searchParams = new URLSearchParams(); if (orderBy) searchParams.append('order_by', orderBy); if (direction) searchParams.append('direction', direction); if (page) searchParams.append('page', page.toString()); + if (modelId) searchParams.append('model_id', modelId); const res = await fetch( `${WEBUI_API_BASE_URL}/evaluations/feedbacks/list?${searchParams.toString()}`, @@ -200,17 +235,23 @@ export const getFeedbackItems = async (token: string = '', orderBy, direction, p return res; }; -export const exportAllFeedbacks = async (token: string = '') => { +export const exportAllFeedbacks = async (token: string = '', modelId: string = '') => { let error = null; - const res = await fetch(`${WEBUI_API_BASE_URL}/evaluations/feedbacks/all/export`, { - method: 'GET', - headers: { - Accept: 'application/json', - 'Content-Type': 'application/json', - authorization: `Bearer ${token}` + const searchParams = new URLSearchParams(); + if (modelId) searchParams.append('model_id', modelId); + + const res = await fetch( + `${WEBUI_API_BASE_URL}/evaluations/feedbacks/all/export?${searchParams.toString()}`, + { + method: 'GET', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + } } - }) + ) .then(async (res) => { if (!res.ok) throw await res.json(); return res.json(); diff --git a/src/lib/apis/index.ts b/src/lib/apis/index.ts index c46c86b801..5faf56d4d4 100644 --- a/src/lib/apis/index.ts +++ b/src/lib/apis/index.ts @@ -161,9 +161,12 @@ export const getModels = async ( type ChatCompletedForm = { model: string; - messages: string[]; + messages: Record[]; chat_id: string; - session_id: string; + session_id: string | undefined; + id: string; + filter_ids?: string[]; + model_item?: unknown; }; export const chatCompleted = async (token: string, body: ChatCompletedForm) => { @@ -270,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', @@ -558,13 +593,24 @@ export const executeToolServer = async ( responseHeaders[key] = value; }); - const text = await res.text(); let responseData; + const contentType = res.headers.get('Content-Type')?.split(';')[0]?.trim() ?? ''; try { - responseData = JSON.parse(text); + responseData = await res.clone().json(); } catch { - responseData = text; + if (contentType.startsWith('text/') || !contentType) { + responseData = await res.text(); + } else { + const buf = await res.arrayBuffer(); + const bytes = new Uint8Array(buf); + let binary = ''; + for (let i = 0; i < bytes.length; i++) { + binary += String.fromCharCode(bytes[i]); + } + const b64 = btoa(binary); + responseData = `data:${contentType};base64,${b64}`; + } } return [responseData, responseHeaders]; } catch (err: any) { diff --git a/src/lib/apis/terminal/index.ts b/src/lib/apis/terminal/index.ts index de2e2fd5a6..69ee2c5a0a 100644 --- a/src/lib/apis/terminal/index.ts +++ b/src/lib/apis/terminal/index.ts @@ -45,7 +45,11 @@ export const getTerminalConfig = async ( return res.json().catch(() => null); }; -export const getCwd = async (baseUrl: string, apiKey: string, sessionId?: string): Promise => { +export const getCwd = async ( + baseUrl: string, + apiKey: string, + sessionId?: string +): Promise => { const url = `${baseUrl.replace(/\/$/, '')}/files/cwd`; const headers: Record = { Authorization: `Bearer ${apiKey}` }; if (sessionId) headers['X-Session-Id'] = sessionId; diff --git a/src/lib/apis/utils/index.ts b/src/lib/apis/utils/index.ts index d19f10f948..e5fea091eb 100644 --- a/src/lib/apis/utils/index.ts +++ b/src/lib/apis/utils/index.ts @@ -16,10 +16,14 @@ export const getGravatarUrl = async (token: string, email: string) => { }) .catch((err) => { console.error(err); - error = err; + error = err.detail ?? err; return null; }); + if (error) { + throw error; + } + return res; }; diff --git a/src/lib/components/AddToolServerModal.svelte b/src/lib/components/AddToolServerModal.svelte index 2237c1afc6..d146aeca0c 100644 --- a/src/lib/components/AddToolServerModal.svelte +++ b/src/lib/components/AddToolServerModal.svelte @@ -917,10 +917,8 @@ 'MCP support is experimental and its specification changes often, which can lead to incompatibilities. OpenAPI specification support is directly maintained by the Open WebUI team, making it the more reliable option for compatibility.' )} - {$i18n.t('Read more →')}{$i18n.t('Read more →')} {/if} diff --git a/src/lib/components/AutomationModal.svelte b/src/lib/components/AutomationModal.svelte new file mode 100644 index 0000000000..c16265515e --- /dev/null +++ b/src/lib/components/AutomationModal.svelte @@ -0,0 +1,161 @@ + + + +
+ +
+ + +
+ + +
+
{$i18n.t('Instructions')}
+