Merge remote-tracking branch 'origin/dev' into fix/retrieval-collection-write-access

This commit is contained in:
DrMelone 2026-04-12 21:53:57 +02:00
commit edf51e5c21
79 changed files with 4983 additions and 4549 deletions

View file

@ -519,6 +519,11 @@ PASSWORD_VALIDATION_HINT = os.environ.get('PASSWORD_VALIDATION_HINT', '')
BYPASS_MODEL_ACCESS_CONTROL = os.environ.get('BYPASS_MODEL_ACCESS_CONTROL', 'False').lower() == 'true'
# When disabled (default), the OpenAI catch-all proxy endpoint (/{path:path})
# is blocked. Enable only if you need direct passthrough to upstream OpenAI-
# compatible APIs for endpoints not natively handled by Open WebUI.
ENABLE_OPENAI_API_PASSTHROUGH = os.environ.get('ENABLE_OPENAI_API_PASSTHROUGH', 'False').lower() == 'true'
WEBUI_AUTH_SIGNOUT_REDIRECT_URL = os.environ.get('WEBUI_AUTH_SIGNOUT_REDIRECT_URL', None)
####################################

View file

@ -53,12 +53,12 @@ logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL)
log = logging.getLogger(__name__)
def get_function_module_by_id(request: Request, pipe_id: str):
function_module, _, _ = get_function_module_from_cache(request, pipe_id)
async def get_function_module_by_id(request: Request, pipe_id: str):
function_module, _, _ = await get_function_module_from_cache(request, pipe_id)
if hasattr(function_module, 'valves') and hasattr(function_module, 'Valves'):
Valves = function_module.Valves
valves = Functions.get_function_valves_by_id(pipe_id)
valves = await Functions.get_function_valves_by_id(pipe_id)
if valves:
try:
@ -73,12 +73,12 @@ def get_function_module_by_id(request: Request, pipe_id: str):
async def get_function_models(request):
pipes = Functions.get_functions_by_type('pipe', active_only=True)
pipes = await Functions.get_functions_by_type('pipe', active_only=True)
pipe_models = []
for pipe in pipes:
try:
function_module = get_function_module_by_id(request, pipe.id)
function_module = await get_function_module_by_id(request, pipe.id)
has_user_valves = False
if hasattr(function_module, 'UserValves'):
@ -187,7 +187,7 @@ async def generate_function_chat_completion(request, form_data, user, models: di
pipe_id, _ = pipe_id.split('.', 1)
return pipe_id
def get_function_params(function_module, form_data, user, extra_params=None):
async def get_function_params(function_module, form_data, user, extra_params=None):
if extra_params is None:
extra_params = {}
@ -198,7 +198,7 @@ async def generate_function_chat_completion(request, form_data, user, models: di
params = {'body': form_data} | {k: v for k, v in extra_params.items() if k in sig.parameters}
if '__user__' in params and hasattr(function_module, 'UserValves'):
user_valves = Functions.get_user_valves_by_id_and_user_id(pipe_id, user.id)
user_valves = await Functions.get_user_valves_by_id_and_user_id(pipe_id, user.id)
try:
params['__user__']['valves'] = function_module.UserValves(**user_valves)
except Exception as e:
@ -208,7 +208,7 @@ async def generate_function_chat_completion(request, form_data, user, models: di
return params
model_id = form_data.get('model')
model_info = Models.get_model_by_id(model_id)
model_info = await Models.get_model_by_id(model_id)
metadata = form_data.pop('metadata', {})
@ -225,8 +225,8 @@ async def generate_function_chat_completion(request, form_data, user, models: di
if metadata:
if all(k in metadata for k in ('session_id', 'chat_id', 'message_id')):
__event_emitter__ = get_event_emitter(metadata)
__event_call__ = get_event_call(metadata)
__event_emitter__ = await get_event_emitter(metadata)
__event_call__ = await get_event_call(metadata)
__task__ = metadata.get('task', None)
__task_body__ = metadata.get('task_body', None)
@ -268,10 +268,10 @@ async def generate_function_chat_completion(request, form_data, user, models: di
form_data = apply_system_prompt_to_body(system, form_data, metadata, user)
pipe_id = get_pipe_id(form_data)
function_module = get_function_module_by_id(request, pipe_id)
function_module = await get_function_module_by_id(request, pipe_id)
pipe = function_module.pipe
params = get_function_params(function_module, form_data, user, extra_params)
params = await get_function_params(function_module, form_data, user, extra_params)
if form_data.get('stream', False):

View file

@ -1,7 +1,7 @@
import os
import json
import logging
from contextlib import contextmanager
from contextlib import asynccontextmanager, contextmanager
from typing import Any, Optional
from open_webui.internal.wrappers import register_connection
@ -19,6 +19,7 @@ from open_webui.env import (
)
from peewee_migrate import Router
from sqlalchemy import Dialect, create_engine, MetaData, event, types
from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker
from sqlalchemy.ext.declarative import declarative_base
from sqlalchemy.orm import scoped_session, sessionmaker, Session
from sqlalchemy.pool import QueuePool, NullPool
@ -81,6 +82,32 @@ if ENABLE_DB_MIGRATIONS:
SQLALCHEMY_DATABASE_URL = DATABASE_URL
def _make_async_url(url: str) -> str:
"""Convert a sync database URL to its async driver equivalent."""
if url.startswith('sqlite+sqlcipher://'):
# SQLCipher has no async driver — not supported for async
raise ValueError(
'sqlite+sqlcipher:// URLs are not supported with async engine. '
'Use standard sqlite:// or postgresql:// instead.'
)
if url.startswith('sqlite:///') or url.startswith('sqlite://'):
return url.replace('sqlite://', 'sqlite+aiosqlite://', 1)
if url.startswith('postgresql+psycopg2://'):
return url.replace('postgresql+psycopg2://', 'postgresql+asyncpg://', 1)
if url.startswith('postgresql://'):
return url.replace('postgresql://', 'postgresql+asyncpg://', 1)
if url.startswith('postgres://'):
return url.replace('postgres://', 'postgresql+asyncpg://', 1)
# For other dialects, return as-is and let SQLAlchemy handle it
return url
# ============================================================
# SYNC ENGINE (used only for: startup migrations, config loading,
# Alembic, peewee migration, health checks)
# ============================================================
# Handle SQLCipher URLs
if SQLALCHEMY_DATABASE_URL.startswith('sqlite+sqlcipher://'):
database_password = os.environ.get('DATABASE_PASSWORD')
@ -155,6 +182,7 @@ else:
engine = create_engine(SQLALCHEMY_DATABASE_URL, pool_pre_ping=True)
# Sync session — used ONLY for startup config loading (config.py runs at import time)
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine, expire_on_commit=False)
metadata_obj = MetaData(schema=DATABASE_SCHEMA)
Base = declarative_base(metadata=metadata_obj)
@ -162,6 +190,7 @@ ScopedSession = scoped_session(SessionLocal)
def get_session():
"""Sync session generator — used ONLY for startup/config operations."""
db = SessionLocal()
try:
yield db
@ -172,10 +201,82 @@ def get_session():
get_db = contextmanager(get_session)
@contextmanager
def get_db_context(db: Optional[Session] = None):
if isinstance(db, Session) and DATABASE_ENABLE_SESSION_SHARING:
# ============================================================
# ASYNC ENGINE (used for ALL runtime database operations)
# ============================================================
ASYNC_SQLALCHEMY_DATABASE_URL = _make_async_url(SQLALCHEMY_DATABASE_URL)
if 'sqlite' in ASYNC_SQLALCHEMY_DATABASE_URL:
async_engine = create_async_engine(
ASYNC_SQLALCHEMY_DATABASE_URL,
connect_args={'check_same_thread': False},
)
if DATABASE_ENABLE_SQLITE_WAL:
@event.listens_for(async_engine.sync_engine, 'connect')
def _set_sqlite_wal(dbapi_connection, connection_record):
cursor = dbapi_connection.cursor()
cursor.execute('PRAGMA journal_mode=WAL')
cursor.close()
else:
if isinstance(DATABASE_POOL_SIZE, int):
if DATABASE_POOL_SIZE > 0:
async_engine = create_async_engine(
ASYNC_SQLALCHEMY_DATABASE_URL,
pool_size=DATABASE_POOL_SIZE,
max_overflow=DATABASE_POOL_MAX_OVERFLOW,
pool_timeout=DATABASE_POOL_TIMEOUT,
pool_recycle=DATABASE_POOL_RECYCLE,
pool_pre_ping=True,
poolclass=QueuePool,
)
else:
async_engine = create_async_engine(
ASYNC_SQLALCHEMY_DATABASE_URL,
pool_pre_ping=True,
poolclass=NullPool,
)
else:
async_engine = create_async_engine(
ASYNC_SQLALCHEMY_DATABASE_URL,
pool_pre_ping=True,
)
AsyncSessionLocal = async_sessionmaker(
bind=async_engine,
class_=AsyncSession,
autocommit=False,
autoflush=False,
expire_on_commit=False,
)
async def get_async_session():
"""Async session generator for FastAPI Depends()."""
async with AsyncSessionLocal() as db:
try:
yield db
finally:
await db.close()
@asynccontextmanager
async def get_async_db():
"""Async context manager for use outside of FastAPI dependency injection."""
async with AsyncSessionLocal() as db:
try:
yield db
finally:
await db.close()
@asynccontextmanager
async def get_async_db_context(db: Optional[AsyncSession] = None):
"""Async context manager that reuses an existing session if provided and session sharing is enabled."""
if isinstance(db, AsyncSession) and DATABASE_ENABLE_SESSION_SHARING:
yield db
else:
with get_db() as session:
async with get_async_db() as session:
yield session

View file

@ -108,8 +108,8 @@ from open_webui.routers.retrieval import (
)
from sqlalchemy.orm import Session
from open_webui.internal.db import ScopedSession, engine, get_session
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import ScopedSession, engine, get_async_session
from open_webui.models.functions import Functions
from open_webui.models.models import Models
@ -575,7 +575,7 @@ from open_webui.constants import ERROR_MESSAGES
if SAFE_MODE:
print('SAFE MODE ENABLED')
Functions.deactivate_all_functions()
# Functions.deactivate_all_functions() is awaited in lifespan below
logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL)
log = logging.getLogger(__name__)
@ -629,14 +629,17 @@ async def lifespan(app: FastAPI):
# Create admin account from env vars if specified and no users exist
if WEBUI_ADMIN_EMAIL and WEBUI_ADMIN_PASSWORD:
if create_admin_user(WEBUI_ADMIN_EMAIL, WEBUI_ADMIN_PASSWORD, WEBUI_ADMIN_NAME):
if await create_admin_user(WEBUI_ADMIN_EMAIL, WEBUI_ADMIN_PASSWORD, WEBUI_ADMIN_NAME):
# Disable signup since we now have an admin
app.state.config.ENABLE_SIGNUP = False
if SAFE_MODE:
await Functions.deactivate_all_functions()
# This should be blocking (sync) so functions are not deactivated on first /get_models calls
# when the first user lands on the / route.
log.info('Installing external dependencies of functions and tools...')
install_tool_and_function_dependencies()
await install_tool_and_function_dependencies()
app.state.redis = get_redis_connection(
redis_url=REDIS_URL,
@ -1605,7 +1608,7 @@ async def get_models(request: Request, refresh: bool = False, user=Depends(get_v
)
)
models = get_filtered_models(models, user)
models = await get_filtered_models(models, user)
log.debug(
f'/api/models returned filtered models accessible to the user: {json.dumps([model.get("id") for model in models])}'
@ -1671,12 +1674,12 @@ async def chat_completion(
raise Exception('Model not found')
model = request.app.state.MODELS[model_id]
model_info = Models.get_model_by_id(model_id)
model_info = await Models.get_model_by_id(model_id)
# Check if user has access to the model
if not BYPASS_MODEL_ACCESS_CONTROL and (user.role != 'admin' or not BYPASS_ADMIN_ACCESS_CONTROL):
try:
check_model_access(user, model)
await check_model_access(user, model)
except Exception as e:
raise e
else:
@ -1758,7 +1761,7 @@ async def chat_completion(
# Verify chat ownership — lightweight EXISTS check avoids
# deserializing the full chat JSON blob just to confirm the row exists
if (
not Chats.is_chat_owner(metadata['chat_id'], user.id) and user.role != 'admin'
not await Chats.is_chat_owner(metadata['chat_id'], user.id) and user.role != 'admin'
): # admins can access any chat
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
@ -1770,7 +1773,7 @@ async def chat_completion(
parent_message_files = parent_message.get('files', [])
if parent_message_files:
try:
Chats.insert_chat_files(
await Chats.insert_chat_files(
metadata['chat_id'],
parent_message.get('id'),
[
@ -1802,7 +1805,7 @@ async def chat_completion(
if metadata.get('chat_id') and metadata.get('message_id'):
try:
if not metadata['chat_id'].startswith('local:'):
Chats.upsert_message_to_chat_by_id_and_message_id(
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
{
@ -1813,13 +1816,13 @@ async def chat_completion(
except Exception:
pass
ctx = build_chat_response_context(request, form_data, user, model, metadata, tasks, events)
ctx = await build_chat_response_context(request, form_data, user, model, metadata, tasks, events)
return await process_chat_response(response, ctx)
except asyncio.CancelledError:
log.info('Chat processing was cancelled')
try:
event_emitter = get_event_emitter(metadata)
event_emitter = await get_event_emitter(metadata)
await asyncio.shield(
event_emitter(
{'type': 'chat:tasks:cancel'},
@ -1835,7 +1838,7 @@ async def chat_completion(
# Update the chat message with the error
try:
if not metadata['chat_id'].startswith('local:'):
Chats.upsert_message_to_chat_by_id_and_message_id(
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
{
@ -1844,7 +1847,7 @@ async def chat_completion(
},
)
event_emitter = get_event_emitter(metadata)
event_emitter = await get_event_emitter(metadata)
await event_emitter(
{
'type': 'chat:message:error',
@ -1883,7 +1886,7 @@ async def chat_completion(
# Emit chat:active=false when task completes
try:
if metadata.get('chat_id'):
event_emitter = get_event_emitter(metadata, update_db=False)
event_emitter = await get_event_emitter(metadata, update_db=False)
if event_emitter:
await event_emitter({'type': 'chat:active', 'data': {'active': False}})
except Exception as e:
@ -1897,7 +1900,7 @@ async def chat_completion(
id=metadata['chat_id'],
)
# Emit chat:active=true when task starts
event_emitter = get_event_emitter(metadata, update_db=False)
event_emitter = await get_event_emitter(metadata, update_db=False)
if event_emitter:
await event_emitter({'type': 'chat:active', 'data': {'active': True}})
return {'status': True, 'task_id': task_id}
@ -2024,7 +2027,7 @@ async def list_tasks_endpoint(request: Request, user=Depends(get_verified_user))
@app.get('/api/tasks/chat/{chat_id}')
async def list_tasks_by_chat_id_endpoint(request: Request, chat_id: str, user=Depends(get_verified_user)):
chat = Chats.get_chat_by_id(chat_id)
chat = await Chats.get_chat_by_id(chat_id)
if chat is None or chat.user_id != user.id:
return {'task_ids': []}
@ -2065,9 +2068,9 @@ async def get_app_config(request: Request):
detail='Invalid token',
)
if data is not None and 'id' in data:
user = Users.get_user_by_id(data['id'])
user = await Users.get_user_by_id(data['id'])
user_count = Users.get_num_users()
user_count = await Users.get_num_users()
onboarding = False
if user is None:
@ -2276,7 +2279,7 @@ async def get_current_usage(user=Depends(get_verified_user)):
return {
'model_ids': get_models_in_use(),
'user_count': Users.get_active_user_count(),
'user_count': await Users.get_active_user_count(),
}
except HTTPException:
raise
@ -2483,7 +2486,7 @@ async def oauth_login_callback(
provider: str,
request: Request,
response: Response,
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
return await oauth_manager.handle_callback(request, provider, response, db=db)
@ -2496,7 +2499,7 @@ async def oauth_login_callback(
@app.post('/oauth/backchannel-logout')
async def oauth_backchannel_logout(
request: Request,
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if not ENABLE_OAUTH_BACKCHANNEL_LOGOUT:
raise HTTPException(status_code=404)

View file

@ -3,8 +3,9 @@ import time
import uuid
from typing import Optional
from sqlalchemy.orm import Session
from open_webui.internal.db import Base, get_db_context
from sqlalchemy import select, delete
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import Base, get_async_db_context
from pydantic import BaseModel, ConfigDict
from sqlalchemy import BigInteger, Column, Text, UniqueConstraint, or_, and_
@ -281,20 +282,20 @@ def grants_to_access_control(grants: list) -> Optional[dict]:
class AccessGrantsTable:
def grant_access(
async def grant_access(
self,
resource_type: str,
resource_id: str,
principal_type: str,
principal_id: str,
permission: str,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[AccessGrantModel]:
"""Add a single access grant. Idempotent (ignores duplicates)."""
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
# Check for existing grant
existing = (
db.query(AccessGrant)
result = await db.execute(
select(AccessGrant)
.filter_by(
resource_type=resource_type,
resource_id=resource_id,
@ -302,8 +303,8 @@ class AccessGrantsTable:
principal_id=principal_id,
permission=permission,
)
.first()
)
existing = result.scalars().first()
if existing:
return AccessGrantModel.model_validate(existing)
@ -317,23 +318,23 @@ class AccessGrantsTable:
created_at=int(time.time()),
)
db.add(grant)
db.commit()
db.refresh(grant)
await db.commit()
await db.refresh(grant)
return AccessGrantModel.model_validate(grant)
def revoke_access(
async def revoke_access(
self,
resource_type: str,
resource_id: str,
principal_type: str,
principal_id: str,
permission: str,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> bool:
"""Remove a single access grant."""
with get_db_context(db) as db:
deleted = (
db.query(AccessGrant)
async with get_async_db_context(db) as db:
result = await db.execute(
delete(AccessGrant)
.filter_by(
resource_type=resource_type,
resource_id=resource_id,
@ -341,47 +342,47 @@ class AccessGrantsTable:
principal_id=principal_id,
permission=permission,
)
.delete()
)
db.commit()
return deleted > 0
await db.commit()
return result.rowcount > 0
def revoke_all_access(
async def revoke_all_access(
self,
resource_type: str,
resource_id: str,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> int:
"""Remove all access grants for a resource."""
with get_db_context(db) as db:
deleted = (
db.query(AccessGrant)
async with get_async_db_context(db) as db:
result = await db.execute(
delete(AccessGrant)
.filter_by(
resource_type=resource_type,
resource_id=resource_id,
)
.delete()
)
db.commit()
return deleted
await db.commit()
return result.rowcount
def set_access_control(
async def set_access_control(
self,
resource_type: str,
resource_id: str,
access_control: Optional[dict],
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> list[AccessGrantModel]:
"""
Replace all grants for a resource from an access_control JSON dict.
This is the primary bridge for backward compat with the frontend.
"""
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
# Delete all existing grants for this resource
db.query(AccessGrant).filter_by(
resource_type=resource_type,
resource_id=resource_id,
).delete()
await db.execute(
delete(AccessGrant).filter_by(
resource_type=resource_type,
resource_id=resource_id,
)
)
# Convert JSON to grant dicts
grant_dicts = access_control_to_grants(resource_type, resource_id, access_control)
@ -397,25 +398,27 @@ class AccessGrantsTable:
db.add(grant)
results.append(grant)
db.commit()
await db.commit()
return [AccessGrantModel.model_validate(g) for g in results]
def set_access_grants(
async def set_access_grants(
self,
resource_type: str,
resource_id: str,
access_grants: Optional[list],
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> list[AccessGrantModel]:
"""
Replace all grants for a resource from a direct access_grants list.
"""
with get_db_context(db) as db:
db.query(AccessGrant).filter_by(
resource_type=resource_type,
resource_id=resource_id,
).delete()
async with get_async_db_context(db) as db:
await db.execute(
delete(AccessGrant).filter_by(
resource_type=resource_type,
resource_id=resource_id,
)
)
normalized_grants = normalize_access_grants(access_grants)
@ -433,80 +436,80 @@ class AccessGrantsTable:
db.add(grant)
results.append(grant)
db.commit()
await db.commit()
return [AccessGrantModel.model_validate(g) for g in results]
def get_access_control(
async def get_access_control(
self,
resource_type: str,
resource_id: str,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[dict]:
"""
Reconstruct the old-style access_control JSON dict from grants.
For backward compat with the frontend.
"""
with get_db_context(db) as db:
grants = (
db.query(AccessGrant)
async with get_async_db_context(db) as db:
result = await db.execute(
select(AccessGrant)
.filter_by(
resource_type=resource_type,
resource_id=resource_id,
)
.all()
)
grants = result.scalars().all()
grant_models = [AccessGrantModel.model_validate(g) for g in grants]
return grants_to_access_control(grant_models)
def get_grants_by_resource(
async def get_grants_by_resource(
self,
resource_type: str,
resource_id: str,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> list[AccessGrantModel]:
"""Get all grants for a specific resource."""
with get_db_context(db) as db:
grants = (
db.query(AccessGrant)
async with get_async_db_context(db) as db:
result = await db.execute(
select(AccessGrant)
.filter_by(
resource_type=resource_type,
resource_id=resource_id,
)
.all()
)
grants = result.scalars().all()
return [AccessGrantModel.model_validate(g) for g in grants]
def get_grants_by_resources(
async def get_grants_by_resources(
self,
resource_type: str,
resource_ids: list[str],
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> dict[str, list[AccessGrantModel]]:
"""Batch-fetch grants for multiple resources. Returns {resource_id: [grants]}."""
if not resource_ids:
return {}
with get_db_context(db) as db:
grants = (
db.query(AccessGrant)
async with get_async_db_context(db) as db:
result = await db.execute(
select(AccessGrant)
.filter(
AccessGrant.resource_type == resource_type,
AccessGrant.resource_id.in_(resource_ids),
)
.all()
)
result: dict[str, list[AccessGrantModel]] = {rid: [] for rid in resource_ids}
grants = result.scalars().all()
result_dict: dict[str, list[AccessGrantModel]] = {rid: [] for rid in resource_ids}
for g in grants:
result[g.resource_id].append(AccessGrantModel.model_validate(g))
return result
result_dict[g.resource_id].append(AccessGrantModel.model_validate(g))
return result_dict
def has_access(
async def has_access(
self,
user_id: str,
resource_type: str,
resource_id: str,
permission: str = 'read',
user_group_ids: Optional[set[str]] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> bool:
"""
Check if a user has the specified permission on a resource.
@ -516,7 +519,7 @@ class AccessGrantsTable:
- There's a grant for the specific user with the requested permission
- There's a grant for any of the user's groups with the requested permission
"""
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
# Build conditions for matching grants
conditions = [
# Public access
@ -535,7 +538,7 @@ class AccessGrantsTable:
if user_group_ids is None:
from open_webui.models.groups import Groups
user_groups = Groups.get_groups_by_member_id(user_id, db=db)
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
user_group_ids = {group.id for group in user_groups}
if user_group_ids:
@ -546,26 +549,27 @@ class AccessGrantsTable:
)
)
exists = (
db.query(AccessGrant)
result = await db.execute(
select(AccessGrant)
.filter(
AccessGrant.resource_type == resource_type,
AccessGrant.resource_id == resource_id,
AccessGrant.permission == permission,
or_(*conditions),
)
.first()
.limit(1)
)
return exists is not None
grant = result.scalars().first()
return grant is not None
def get_accessible_resource_ids(
async def get_accessible_resource_ids(
self,
user_id: str,
resource_type: str,
resource_ids: list[str],
permission: str = 'read',
user_group_ids: Optional[set[str]] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> set[str]:
"""
Batch check: return the subset of resource_ids that the user can access.
@ -575,7 +579,7 @@ class AccessGrantsTable:
if not resource_ids:
return set()
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
conditions = [
and_(
AccessGrant.principal_type == 'user',
@ -590,7 +594,7 @@ class AccessGrantsTable:
if user_group_ids is None:
from open_webui.models.groups import Groups
user_groups = Groups.get_groups_by_member_id(user_id, db=db)
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
user_group_ids = {group.id for group in user_groups}
if user_group_ids:
@ -601,8 +605,8 @@ class AccessGrantsTable:
)
)
rows = (
db.query(AccessGrant.resource_id)
result = await db.execute(
select(AccessGrant.resource_id)
.filter(
AccessGrant.resource_type == resource_type,
AccessGrant.resource_id.in_(resource_ids),
@ -610,16 +614,16 @@ class AccessGrantsTable:
or_(*conditions),
)
.distinct()
.all()
)
rows = result.all()
return {row[0] for row in rows}
def get_users_with_access(
async def get_users_with_access(
self,
resource_type: str,
resource_id: str,
permission: str = 'read',
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> list:
"""
Get all users who have the specified permission on a resource.
@ -628,21 +632,21 @@ class AccessGrantsTable:
from open_webui.models.users import Users, UserModel
from open_webui.models.groups import Groups
with get_db_context(db) as db:
grants = (
db.query(AccessGrant)
async with get_async_db_context(db) as db:
result = await db.execute(
select(AccessGrant)
.filter_by(
resource_type=resource_type,
resource_id=resource_id,
permission=permission,
)
.all()
)
grants = result.scalars().all()
# Check for public access
for grant in grants:
if grant.principal_type == 'user' and grant.principal_id == '*':
result = Users.get_users(filter={'roles': ['!pending']}, db=db)
result = await Users.get_users(filter={'roles': ['!pending']}, db=db)
return result.get('users', [])
user_ids_with_access = set()
@ -651,14 +655,14 @@ class AccessGrantsTable:
if grant.principal_type == 'user':
user_ids_with_access.add(grant.principal_id)
elif grant.principal_type == 'group':
group_user_ids = Groups.get_group_user_ids_by_id(grant.principal_id, db=db)
group_user_ids = await Groups.get_group_user_ids_by_id(grant.principal_id, db=db)
if group_user_ids:
user_ids_with_access.update(group_user_ids)
if not user_ids_with_access:
return []
return Users.get_users_by_user_ids(list(user_ids_with_access), db=db)
return await Users.get_users_by_user_ids(list(user_ids_with_access), db=db)
def has_permission_filter(
self,
@ -673,6 +677,10 @@ class AccessGrantsTable:
Apply access control filtering to a SQLAlchemy query by JOINing with access_grant.
This replaces the old JSON-column-based filtering with a proper relational JOIN.
Note: This method builds SQLAlchemy expressions and does NOT perform I/O itself,
so it remains synchronous. The caller is responsible for executing the query
asynchronously with `await db.execute(...)`.
"""
group_ids = filter.get('group_ids', [])
user_id = filter.get('user_id')
@ -718,7 +726,7 @@ class AccessGrantsTable:
# LEFT JOIN access_grant and filter
# We use a subquery approach to avoid duplicates from multiple matching grants
from sqlalchemy import exists as sa_exists, select
from sqlalchemy import exists as sa_exists
grant_exists = (
select(AccessGrant.id)
@ -776,11 +784,15 @@ class AccessGrantsTable:
"""
Filter for items where user has read BUT NOT write access.
Public items are NOT considered read_only.
Note: This method builds SQLAlchemy expressions and does NOT perform I/O itself,
so it remains synchronous. The caller is responsible for executing the query
asynchronously with `await db.execute(...)`.
"""
group_ids = filter.get('group_ids', [])
user_id = filter.get('user_id')
from sqlalchemy import exists as sa_exists, select
from sqlalchemy import exists as sa_exists
# Has read grant (not public)
read_grant_exists = (

View file

@ -2,8 +2,9 @@ import logging
import uuid
from typing import Optional
from sqlalchemy.orm import Session
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
from sqlalchemy import select, delete, update
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import Base, JSONField, get_async_db_context
from open_webui.models.users import User, UserModel, UserProfileImageResponse, Users
from open_webui.utils.validate import validate_profile_image_url
from pydantic import BaseModel, field_validator
@ -88,7 +89,7 @@ class AddUserForm(SignupForm):
class AuthsTable:
def insert_new_auth(
async def insert_new_auth(
self,
email: str,
password: str,
@ -96,9 +97,9 @@ class AuthsTable:
profile_image_url: str = '/user.png',
role: str = 'pending',
oauth: Optional[dict] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[UserModel]:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
log.info('insert_new_auth')
id = str(uuid.uuid4())
@ -107,28 +108,29 @@ class AuthsTable:
result = Auth(**auth.model_dump())
db.add(result)
user = Users.insert_new_user(id, name, email, profile_image_url, role, oauth=oauth, db=db)
user = await Users.insert_new_user(id, name, email, profile_image_url, role, oauth=oauth, db=db)
db.commit()
db.refresh(result)
await db.commit()
await db.refresh(result)
if result and user:
return user
else:
return None
def authenticate_user(
self, email: str, verify_password: callable, db: Optional[Session] = None
async def authenticate_user(
self, email: str, verify_password: callable, db: Optional[AsyncSession] = None
) -> Optional[UserModel]:
log.info(f'authenticate_user: {email}')
user = Users.get_user_by_email(email, db=db)
user = await Users.get_user_by_email(email, db=db)
if not user:
return None
try:
with get_db_context(db) as db:
auth = db.query(Auth).filter_by(id=user.id, active=True).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(Auth).filter_by(id=user.id, active=True))
auth = result.scalars().first()
if auth:
if verify_password(auth.password):
return user
@ -139,66 +141,66 @@ class AuthsTable:
except Exception:
return None
def authenticate_user_by_api_key(self, api_key: str, db: Optional[Session] = None) -> Optional[UserModel]:
async def authenticate_user_by_api_key(self, api_key: str, db: Optional[AsyncSession] = None) -> Optional[UserModel]:
log.info(f'authenticate_user_by_api_key')
# if no api_key, return None
if not api_key:
return None
try:
user = Users.get_user_by_api_key(api_key, db=db)
user = await Users.get_user_by_api_key(api_key, db=db)
return user if user else None
except Exception:
return False
def authenticate_user_by_email(self, email: str, db: Optional[Session] = None) -> Optional[UserModel]:
async def authenticate_user_by_email(self, email: str, db: Optional[AsyncSession] = None) -> Optional[UserModel]:
log.info(f'authenticate_user_by_email: {email}')
try:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
# Single JOIN query instead of two separate queries
result = (
db.query(Auth, User)
result = await db.execute(
select(Auth, User)
.join(User, Auth.id == User.id)
.filter(Auth.email == email, Auth.active == True)
.first()
)
if result:
_, user = result
row = result.first()
if row:
_, user = row
return UserModel.model_validate(user)
return None
except Exception:
return None
def update_user_password_by_id(self, id: str, new_password: str, db: Optional[Session] = None) -> bool:
async def update_user_password_by_id(self, id: str, new_password: str, db: Optional[AsyncSession] = None) -> bool:
try:
with get_db_context(db) as db:
result = db.query(Auth).filter_by(id=id).update({'password': new_password})
db.commit()
return True if result == 1 else False
async with get_async_db_context(db) as db:
result = await db.execute(update(Auth).filter_by(id=id).values(password=new_password))
await db.commit()
return True if result.rowcount == 1 else False
except Exception:
return False
def update_email_by_id(self, id: str, email: str, db: Optional[Session] = None) -> bool:
async def update_email_by_id(self, id: str, email: str, db: Optional[AsyncSession] = None) -> bool:
try:
with get_db_context(db) as db:
result = db.query(Auth).filter_by(id=id).update({'email': email})
db.commit()
if result == 1:
Users.update_user_by_id(id, {'email': email}, db=db)
async with get_async_db_context(db) as db:
result = await db.execute(update(Auth).filter_by(id=id).values(email=email))
await db.commit()
if result.rowcount == 1:
await Users.update_user_by_id(id, {'email': email}, db=db)
return True
return False
except Exception:
return False
def delete_auth_by_id(self, id: str, db: Optional[Session] = None) -> bool:
async def delete_auth_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
try:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
# Delete User
result = Users.delete_user_by_id(id, db=db)
result = await Users.delete_user_by_id(id, db=db)
if result:
db.query(Auth).filter_by(id=id).delete()
db.commit()
await db.execute(delete(Auth).filter_by(id=id))
await db.commit()
return True
else:

View file

@ -4,10 +4,10 @@ from typing import Optional
from uuid import uuid4
from pydantic import BaseModel, ConfigDict
from sqlalchemy import Column, Text, JSON, Boolean, BigInteger, Index, select, or_, func, cast, String
from sqlalchemy.orm import Session
from sqlalchemy import Column, Text, JSON, Boolean, BigInteger, Index, select, or_, func, cast, String, delete, update
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import Base, get_db, get_db_context
from open_webui.internal.db import Base, get_async_db_context
log = logging.getLogger(__name__)
@ -118,14 +118,14 @@ class AutomationListResponse(BaseModel):
class AutomationTable:
def insert(
async def insert(
self,
user_id: str,
form: AutomationForm,
next_run_at: int,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> AutomationModel:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
now = int(time.time_ns())
row = Automation(
id=str(uuid4()),
@ -139,35 +139,38 @@ class AutomationTable:
updated_at=now,
)
db.add(row)
db.commit()
db.refresh(row)
await db.commit()
await db.refresh(row)
return AutomationModel.model_validate(row)
def count_by_user(self, user_id: str, db: Optional[Session] = None) -> int:
with get_db_context(db) as db:
return db.query(Automation).filter_by(user_id=user_id).count()
async def count_by_user(self, user_id: str, db: Optional[AsyncSession] = None) -> int:
async with get_async_db_context(db) as db:
result = await db.execute(
select(func.count()).select_from(Automation).filter_by(user_id=user_id)
)
return result.scalar()
def get_by_id(self, id: str, db: Optional[Session] = None) -> Optional[AutomationModel]:
with get_db_context(db) as db:
row = db.get(Automation, id)
async def get_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[AutomationModel]:
async with get_async_db_context(db) as db:
row = await db.get(Automation, id)
return AutomationModel.model_validate(row) if row else None
def search_automations(
async def search_automations(
self,
user_id: str,
query: Optional[str] = None,
status: Optional[str] = None,
skip: int = 0,
limit: int = 30,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> 'AutomationListResponse':
with get_db_context(db) as db:
q = db.query(Automation).filter_by(user_id=user_id)
async with get_async_db_context(db) as db:
stmt = select(Automation).filter_by(user_id=user_id)
if query:
search = f'%{query}%'
# Search in name and prompt inside JSON data
q = q.filter(
stmt = stmt.filter(
or_(
Automation.name.ilike(search),
cast(Automation.data, String).ilike(search),
@ -175,34 +178,39 @@ class AutomationTable:
)
if status == 'active':
q = q.filter(Automation.is_active == True)
stmt = stmt.filter(Automation.is_active == True)
elif status == 'paused':
q = q.filter(Automation.is_active == False)
stmt = stmt.filter(Automation.is_active == False)
q = q.order_by(Automation.created_at.desc())
stmt = stmt.order_by(Automation.created_at.desc())
total = q.count()
# Get total count
count_result = await db.execute(
select(func.count()).select_from(stmt.subquery())
)
total = count_result.scalar()
if skip:
q = q.offset(skip)
stmt = stmt.offset(skip)
if limit:
q = q.limit(limit)
stmt = stmt.limit(limit)
rows = q.all()
result = await db.execute(stmt)
rows = result.scalars().all()
return AutomationListResponse(
items=[AutomationModel.model_validate(r) for r in rows],
total=total,
)
def update_by_id(
async def update_by_id(
self,
id: str,
form: AutomationForm,
next_run_at: int,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[AutomationModel]:
with get_db_context(db) as db:
row = db.get(Automation, id)
async with get_async_db_context(db) as db:
row = await db.get(Automation, id)
if not row:
return None
row.name = form.name
@ -212,37 +220,37 @@ class AutomationTable:
row.is_active = form.is_active
row.next_run_at = next_run_at
row.updated_at = int(time.time_ns())
db.commit()
db.refresh(row)
await db.commit()
await db.refresh(row)
return AutomationModel.model_validate(row)
def toggle(
async def toggle(
self,
id: str,
next_run_at: Optional[int],
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[AutomationModel]:
with get_db_context(db) as db:
row = db.get(Automation, id)
async with get_async_db_context(db) as db:
row = await db.get(Automation, id)
if not row:
return None
row.is_active = not row.is_active
row.next_run_at = next_run_at if row.is_active else None
row.updated_at = int(time.time_ns())
db.commit()
db.refresh(row)
await db.commit()
await db.refresh(row)
return AutomationModel.model_validate(row)
def delete(self, id: str, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
row = db.get(Automation, id)
async def delete(self, id: str, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
row = await db.get(Automation, id)
if not row:
return False
db.delete(row)
db.commit()
await db.delete(row)
await db.commit()
return True
def claim_due(self, now_ns: int, limit: int = 10, db: Optional[Session] = None) -> list[AutomationModel]:
async def claim_due(self, now_ns: int, limit: int = 10, db: Optional[AsyncSession] = None) -> list[AutomationModel]:
"""
Atomically claim due automations for execution.
@ -250,7 +258,7 @@ class AutomationTable:
double-claimed. On PostgreSQL, uses FOR UPDATE SKIP LOCKED
for zero-contention distributed work claiming.
"""
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
stmt = (
select(Automation)
.where(
@ -264,7 +272,8 @@ class AutomationTable:
if db.bind.dialect.name == 'postgresql':
stmt = stmt.with_for_update(skip_locked=True)
rows = db.execute(stmt).scalars().all()
result = await db.execute(stmt)
rows = result.scalars().all()
from open_webui.utils.automations import next_run_ns
@ -272,7 +281,7 @@ class AutomationTable:
row.last_run_at = now_ns
row.next_run_at = next_run_ns(row.data.get('rrule', ''))
db.commit()
await db.commit()
return [AutomationModel.model_validate(r) for r in rows]
@ -283,15 +292,15 @@ class AutomationTable:
class AutomationRunTable:
def insert(
async def insert(
self,
automation_id: str,
status: str,
chat_id: Optional[str] = None,
error: Optional[str] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> AutomationRunModel:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
row = AutomationRun(
id=str(uuid4()),
automation_id=automation_id,
@ -301,30 +310,31 @@ class AutomationRunTable:
created_at=int(time.time_ns()),
)
db.add(row)
db.commit()
db.refresh(row)
await db.commit()
await db.refresh(row)
return AutomationRunModel.model_validate(row)
def get_latest(self, automation_id: str, db: Optional[Session] = None) -> Optional[AutomationRunModel]:
with get_db_context(db) as db:
row = (
db.query(AutomationRun)
async def get_latest(self, automation_id: str, db: Optional[AsyncSession] = None) -> Optional[AutomationRunModel]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(AutomationRun)
.filter_by(automation_id=automation_id)
.order_by(AutomationRun.created_at.desc())
.first()
.limit(1)
)
row = result.scalars().first()
return AutomationRunModel.model_validate(row) if row else None
def get_latest_batch(
self, automation_ids: list[str], db: Optional[Session] = None
async def get_latest_batch(
self, automation_ids: list[str], db: Optional[AsyncSession] = None
) -> dict[str, AutomationRunModel]:
"""Fetch the latest run for each automation in a single query."""
if not automation_ids:
return {}
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
# Subquery: max created_at per automation_id
subq = (
db.query(
select(
AutomationRun.automation_id,
func.max(AutomationRun.created_at).label('max_created'),
)
@ -332,43 +342,43 @@ class AutomationRunTable:
.group_by(AutomationRun.automation_id)
.subquery()
)
rows = (
db.query(AutomationRun)
result = await db.execute(
select(AutomationRun)
.join(
subq,
(AutomationRun.automation_id == subq.c.automation_id)
& (AutomationRun.created_at == subq.c.max_created),
)
.all()
)
rows = result.scalars().all()
return {
row.automation_id: AutomationRunModel.model_validate(row)
for row in rows
}
def get_by_automation(
async def get_by_automation(
self,
automation_id: str,
skip: int = 0,
limit: int = 50,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> list[AutomationRunModel]:
with get_db_context(db) as db:
rows = (
db.query(AutomationRun)
async with get_async_db_context(db) as db:
result = await db.execute(
select(AutomationRun)
.filter_by(automation_id=automation_id)
.order_by(AutomationRun.created_at.desc())
.offset(skip)
.limit(limit)
.all()
)
rows = result.scalars().all()
return [AutomationRunModel.model_validate(r) for r in rows]
def delete_by_automation(self, automation_id: str, db: Optional[Session] = None) -> int:
with get_db_context(db) as db:
count = db.query(AutomationRun).filter_by(automation_id=automation_id).delete()
db.commit()
return count
async def delete_by_automation(self, automation_id: str, db: Optional[AsyncSession] = None) -> int:
async with get_async_db_context(db) as db:
result = await db.execute(delete(AutomationRun).filter_by(automation_id=automation_id))
await db.commit()
return result.rowcount
Automations = AutomationTable()

View file

@ -4,8 +4,9 @@ import time
import uuid
from typing import Optional
from sqlalchemy.orm import Session
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
from sqlalchemy import select, delete, update, func, case, or_, and_
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import Base, JSONField, get_async_db_context
from open_webui.models.groups import Groups
from open_webui.models.access_grants import (
AccessGrantModel,
@ -25,11 +26,7 @@ from sqlalchemy import (
Text,
JSON,
UniqueConstraint,
case,
cast,
)
from sqlalchemy import or_, func, select, and_, text
from sqlalchemy.sql import exists
####################
# Channel DB Schema
@ -249,22 +246,22 @@ class ChannelWebhookForm(BaseModel):
class ChannelTable:
def _get_access_grants(self, channel_id: str, db: Optional[Session] = None) -> list[AccessGrantModel]:
return AccessGrants.get_grants_by_resource('channel', channel_id, db=db)
async def _get_access_grants(self, channel_id: str, db: Optional[AsyncSession] = None) -> list[AccessGrantModel]:
return await AccessGrants.get_grants_by_resource('channel', channel_id, db=db)
def _to_channel_model(
async def _to_channel_model(
self,
channel: Channel,
access_grants: Optional[list[AccessGrantModel]] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> ChannelModel:
channel_data = ChannelModel.model_validate(channel).model_dump(exclude={'access_grants'})
channel_data['access_grants'] = (
access_grants if access_grants is not None else self._get_access_grants(channel_data['id'], db=db)
access_grants if access_grants is not None else await self._get_access_grants(channel_data['id'], db=db)
)
return ChannelModel.model_validate(channel_data)
def _collect_unique_user_ids(
async def _collect_unique_user_ids(
self,
invited_by: str,
user_ids: Optional[list[str]] = None,
@ -281,7 +278,8 @@ class ChannelTable:
users.add(invited_by)
for group_id in group_ids or []:
users.update(Groups.get_group_user_ids_by_id(group_id))
group_user_ids = await Groups.get_group_user_ids_by_id(group_id)
users.update(group_user_ids)
return users
@ -321,10 +319,20 @@ class ChannelTable:
return memberships
def insert_new_channel(
self, form_data: CreateChannelForm, user_id: str, db: Optional[Session] = None
def _has_permission(self, db, query, filter: dict, permission: str = 'read'):
return AccessGrants.has_permission_filter(
db=db,
query=query,
DocumentModel=Channel,
filter=filter,
resource_type='channel',
permission=permission,
)
async def insert_new_channel(
self, form_data: CreateChannelForm, user_id: str, db: Optional[AsyncSession] = None
) -> Optional[ChannelModel]:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
channel = ChannelModel(
**{
**form_data.model_dump(exclude={'access_grants'}),
@ -340,7 +348,7 @@ class ChannelTable:
new_channel = Channel(**channel.model_dump(exclude={'access_grants'}))
if form_data.type in ['group', 'dm']:
users = self._collect_unique_user_ids(
users = await self._collect_unique_user_ids(
invited_by=user_id,
user_ids=form_data.user_ids,
group_ids=form_data.group_ids,
@ -353,17 +361,18 @@ class ChannelTable:
db.add_all(memberships)
db.add(new_channel)
db.commit()
AccessGrants.set_access_grants('channel', new_channel.id, form_data.access_grants, db=db)
return self._to_channel_model(new_channel, db=db)
await db.commit()
await AccessGrants.set_access_grants('channel', new_channel.id, form_data.access_grants, db=db)
return await self._to_channel_model(new_channel, db=db)
def get_channels(self, db: Optional[Session] = None) -> list[ChannelModel]:
with get_db_context(db) as db:
channels = db.query(Channel).all()
async def get_channels(self, db: Optional[AsyncSession] = None) -> list[ChannelModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(Channel))
channels = result.scalars().all()
channel_ids = [channel.id for channel in channels]
grants_map = AccessGrants.get_grants_by_resources('channel', channel_ids, db=db)
grants_map = await AccessGrants.get_grants_by_resources('channel', channel_ids, db=db)
return [
self._to_channel_model(
await self._to_channel_model(
channel,
access_grants=grants_map.get(channel.id, []),
db=db,
@ -371,22 +380,12 @@ class ChannelTable:
for channel in channels
]
def _has_permission(self, db, query, filter: dict, permission: str = 'read'):
return AccessGrants.has_permission_filter(
db=db,
query=query,
DocumentModel=Channel,
filter=filter,
resource_type='channel',
permission=permission,
)
async def get_channels_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[ChannelModel]:
async with get_async_db_context(db) as db:
user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id, db=db)]
def get_channels_by_user_id(self, user_id: str, db: Optional[Session] = None) -> list[ChannelModel]:
with get_db_context(db) as db:
user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id, db=db)]
membership_channels = (
db.query(Channel)
result = await db.execute(
select(Channel)
.join(ChannelMember, Channel.id == ChannelMember.channel_id)
.filter(
Channel.deleted_at.is_(None),
@ -395,10 +394,10 @@ class ChannelTable:
ChannelMember.user_id == user_id,
ChannelMember.is_active.is_(True),
)
.all()
)
membership_channels = result.scalars().all()
query = db.query(Channel).filter(
stmt = select(Channel).filter(
Channel.deleted_at.is_(None),
Channel.archived_at.is_(None),
or_(
@ -407,17 +406,18 @@ class ChannelTable:
and_(Channel.type != 'group', Channel.type != 'dm'),
),
)
query = self._has_permission(db, query, {'user_id': user_id, 'group_ids': user_group_ids})
stmt = self._has_permission(db, stmt, {'user_id': user_id, 'group_ids': user_group_ids})
standard_channels = query.all()
result = await db.execute(stmt)
standard_channels = result.scalars().all()
all_channels = membership_channels + standard_channels
all_channels = list(membership_channels) + list(standard_channels)
channel_ids = [c.id for c in all_channels]
grants_map = AccessGrants.get_grants_by_resources('channel', channel_ids, db=db)
return [self._to_channel_model(c, access_grants=grants_map.get(c.id, []), db=db) for c in all_channels]
grants_map = await AccessGrants.get_grants_by_resources('channel', channel_ids, db=db)
return [await self._to_channel_model(c, access_grants=grants_map.get(c.id, []), db=db) for c in all_channels]
def get_dm_channel_by_user_ids(self, user_ids: list[str], db: Optional[Session] = None) -> Optional[ChannelModel]:
with get_db_context(db) as db:
async def get_dm_channel_by_user_ids(self, user_ids: list[str], db: Optional[AsyncSession] = None) -> Optional[ChannelModel]:
async with get_async_db_context(db) as db:
# Ensure uniqueness in case a list with duplicates is passed
unique_user_ids = list(set(user_ids))
@ -429,7 +429,7 @@ class ChannelTable:
)
subquery = (
db.query(ChannelMember.channel_id)
select(ChannelMember.channel_id)
.group_by(ChannelMember.channel_id)
# 1. Channel must have exactly len(user_ids) members
.having(func.count(ChannelMember.user_id) == len(unique_user_ids))
@ -438,33 +438,34 @@ class ChannelTable:
.subquery()
)
channel = (
db.query(Channel)
result = await db.execute(
select(Channel)
.filter(
Channel.id.in_(subquery),
Channel.id.in_(select(subquery.c.channel_id)),
Channel.type == 'dm',
)
.first()
.limit(1)
)
channel = result.scalars().first()
return self._to_channel_model(channel, db=db) if channel else None
return await self._to_channel_model(channel, db=db) if channel else None
def add_members_to_channel(
async def add_members_to_channel(
self,
channel_id: str,
invited_by: str,
user_ids: Optional[list[str]] = None,
group_ids: Optional[list[str]] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> list[ChannelMemberModel]:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
# 1. Collect all user_ids including groups + inviter
requested_users = self._collect_unique_user_ids(invited_by, user_ids, group_ids)
requested_users = await self._collect_unique_user_ids(invited_by, user_ids, group_ids)
existing_users = {
row.user_id
for row in db.query(ChannelMember.user_id).filter(ChannelMember.channel_id == channel_id).all()
}
result = await db.execute(
select(ChannelMember.user_id).filter(ChannelMember.channel_id == channel_id)
)
existing_users = {row[0] for row in result.all()}
new_user_ids = requested_users - existing_users
if not new_user_ids:
@ -473,58 +474,54 @@ class ChannelTable:
new_memberships = self._create_membership_models(channel_id, invited_by, new_user_ids)
db.add_all(new_memberships)
db.commit()
await db.commit()
return [ChannelMemberModel.model_validate(membership) for membership in new_memberships]
def remove_members_from_channel(
async def remove_members_from_channel(
self,
channel_id: str,
user_ids: list[str],
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> int:
with get_db_context(db) as db:
result = (
db.query(ChannelMember)
.filter(
async with get_async_db_context(db) as db:
result = await db.execute(
delete(ChannelMember).filter(
ChannelMember.channel_id == channel_id,
ChannelMember.user_id.in_(user_ids),
)
.delete(synchronize_session=False)
)
db.commit()
return result # number of rows deleted
await db.commit()
return result.rowcount # number of rows deleted
def is_user_channel_manager(self, channel_id: str, user_id: str, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
# Check if the user is the creator of the channel
# or has a 'manager' role in ChannelMember
channel = db.query(Channel).filter(Channel.id == channel_id).first()
async def is_user_channel_manager(self, channel_id: str, user_id: str, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
result = await db.execute(select(Channel).filter(Channel.id == channel_id))
channel = result.scalars().first()
if channel and channel.user_id == user_id:
return True
membership = (
db.query(ChannelMember)
.filter(
result = await db.execute(
select(ChannelMember).filter(
ChannelMember.channel_id == channel_id,
ChannelMember.user_id == user_id,
ChannelMember.is_active.is_(True),
ChannelMember.role == 'manager',
)
.first()
)
membership = result.scalars().first()
return membership is not None
def join_channel(self, channel_id: str, user_id: str, db: Optional[Session] = None) -> Optional[ChannelMemberModel]:
with get_db_context(db) as db:
async def join_channel(self, channel_id: str, user_id: str, db: Optional[AsyncSession] = None) -> Optional[ChannelMemberModel]:
async with get_async_db_context(db) as db:
# Check if the membership already exists
existing_membership = (
db.query(ChannelMember)
.filter(
result = await db.execute(
select(ChannelMember).filter(
ChannelMember.channel_id == channel_id,
ChannelMember.user_id == user_id,
)
.first()
)
existing_membership = result.scalars().first()
if existing_membership:
return ChannelMemberModel.model_validate(existing_membership)
@ -548,19 +545,18 @@ class ChannelTable:
new_membership = ChannelMember(**channel_member.model_dump())
db.add(new_membership)
db.commit()
await db.commit()
return channel_member
def leave_channel(self, channel_id: str, user_id: str, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
membership = (
db.query(ChannelMember)
.filter(
async def leave_channel(self, channel_id: str, user_id: str, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
result = await db.execute(
select(ChannelMember).filter(
ChannelMember.channel_id == channel_id,
ChannelMember.user_id == user_id,
)
.first()
)
membership = result.scalars().first()
if not membership:
return False
@ -569,125 +565,127 @@ class ChannelTable:
membership.left_at = int(time.time_ns())
membership.updated_at = int(time.time_ns())
db.commit()
await db.commit()
return True
def get_member_by_channel_and_user_id(
self, channel_id: str, user_id: str, db: Optional[Session] = None
async def get_member_by_channel_and_user_id(
self, channel_id: str, user_id: str, db: Optional[AsyncSession] = None
) -> Optional[ChannelMemberModel]:
with get_db_context(db) as db:
membership = (
db.query(ChannelMember)
.filter(
async with get_async_db_context(db) as db:
result = await db.execute(
select(ChannelMember).filter(
ChannelMember.channel_id == channel_id,
ChannelMember.user_id == user_id,
)
.first()
)
membership = result.scalars().first()
return ChannelMemberModel.model_validate(membership) if membership else None
def get_members_by_channel_id(self, channel_id: str, db: Optional[Session] = None) -> list[ChannelMemberModel]:
with get_db_context(db) as db:
memberships = db.query(ChannelMember).filter(ChannelMember.channel_id == channel_id).all()
async def get_members_by_channel_id(self, channel_id: str, db: Optional[AsyncSession] = None) -> list[ChannelMemberModel]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(ChannelMember).filter(ChannelMember.channel_id == channel_id)
)
memberships = result.scalars().all()
return [ChannelMemberModel.model_validate(membership) for membership in memberships]
def pin_channel(
async def pin_channel(
self,
channel_id: str,
user_id: str,
is_pinned: bool,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> bool:
with get_db_context(db) as db:
membership = (
db.query(ChannelMember)
.filter(
async with get_async_db_context(db) as db:
result = await db.execute(
select(ChannelMember).filter(
ChannelMember.channel_id == channel_id,
ChannelMember.user_id == user_id,
)
.first()
)
membership = result.scalars().first()
if not membership:
return False
membership.is_channel_pinned = is_pinned
membership.updated_at = int(time.time_ns())
db.commit()
await db.commit()
return True
def update_member_last_read_at(self, channel_id: str, user_id: str, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
membership = (
db.query(ChannelMember)
.filter(
async def update_member_last_read_at(self, channel_id: str, user_id: str, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
result = await db.execute(
select(ChannelMember).filter(
ChannelMember.channel_id == channel_id,
ChannelMember.user_id == user_id,
)
.first()
)
membership = result.scalars().first()
if not membership:
return False
membership.last_read_at = int(time.time_ns())
membership.updated_at = int(time.time_ns())
db.commit()
await db.commit()
return True
def update_member_active_status(
async def update_member_active_status(
self,
channel_id: str,
user_id: str,
is_active: bool,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> bool:
with get_db_context(db) as db:
membership = (
db.query(ChannelMember)
.filter(
async with get_async_db_context(db) as db:
result = await db.execute(
select(ChannelMember).filter(
ChannelMember.channel_id == channel_id,
ChannelMember.user_id == user_id,
)
.first()
)
membership = result.scalars().first()
if not membership:
return False
membership.is_active = is_active
membership.updated_at = int(time.time_ns())
db.commit()
await db.commit()
return True
def is_user_channel_member(self, channel_id: str, user_id: str, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
membership = (
db.query(ChannelMember)
.filter(
async def is_user_channel_member(self, channel_id: str, user_id: str, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
result = await db.execute(
select(ChannelMember).filter(
ChannelMember.channel_id == channel_id,
ChannelMember.user_id == user_id,
)
.first()
ChannelMember.is_active.is_(True),
).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 +693,123 @@ class ChannelTable:
for channel in channels
]
def get_channels_by_file_id_and_user_id(
self, file_id: str, user_id: str, db: Optional[Session] = None
async def get_channels_by_file_id_and_user_id(
self, file_id: str, user_id: str, db: Optional[AsyncSession] = None
) -> list[ChannelModel]:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
# 1. Determine which channels have this file
channel_file_rows = db.query(ChannelFile).filter(ChannelFile.file_id == file_id).all()
result = await db.execute(select(ChannelFile).filter(ChannelFile.file_id == file_id))
channel_file_rows = result.scalars().all()
channel_ids = [row.channel_id for row in channel_file_rows]
if not channel_ids:
return []
# 2. Load all channel rows that still exist
channels = (
db.query(Channel)
.filter(
result = await db.execute(
select(Channel).filter(
Channel.id.in_(channel_ids),
Channel.deleted_at.is_(None),
Channel.archived_at.is_(None),
)
.all()
)
channels = result.scalars().all()
if not channels:
return []
# Preload user's group membership
user_group_ids = [g.id for g in Groups.get_groups_by_member_id(user_id, db=db)]
user_group_ids = [g.id for g in await Groups.get_groups_by_member_id(user_id, db=db)]
allowed_channels = []
for channel in channels:
# --- Case A: group or dm => user must be an active member ---
if channel.type in ['group', 'dm']:
membership = (
db.query(ChannelMember)
.filter(
result = await db.execute(
select(ChannelMember).filter(
ChannelMember.channel_id == channel.id,
ChannelMember.user_id == user_id,
ChannelMember.is_active.is_(True),
)
.first()
).limit(1)
)
membership = result.scalars().first()
if membership:
allowed_channels.append(self._to_channel_model(channel, db=db))
allowed_channels.append(await self._to_channel_model(channel, db=db))
continue
# --- Case B: standard channel => rely on ACL permissions ---
query = db.query(Channel).filter(Channel.id == channel.id)
stmt = select(Channel).filter(Channel.id == channel.id)
query = self._has_permission(
stmt = self._has_permission(
db,
query,
stmt,
{'user_id': user_id, 'group_ids': user_group_ids},
permission='read',
)
allowed = query.first()
result = await db.execute(stmt)
allowed = result.scalars().first()
if allowed:
allowed_channels.append(self._to_channel_model(allowed, db=db))
allowed_channels.append(await self._to_channel_model(allowed, db=db))
return allowed_channels
def get_channel_by_id_and_user_id(
self, id: str, user_id: str, db: Optional[Session] = None
async def get_channel_by_id_and_user_id(
self, id: str, user_id: str, db: Optional[AsyncSession] = None
) -> Optional[ChannelModel]:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
# Fetch the channel
channel: Channel = (
db.query(Channel)
.filter(
result = await db.execute(
select(Channel).filter(
Channel.id == id,
Channel.deleted_at.is_(None),
Channel.archived_at.is_(None),
)
.first()
)
channel = result.scalars().first()
if not channel:
return None
# If the channel is a group or dm, read access requires membership (active)
if channel.type in ['group', 'dm']:
membership = (
db.query(ChannelMember)
.filter(
result = await db.execute(
select(ChannelMember).filter(
ChannelMember.channel_id == id,
ChannelMember.user_id == user_id,
ChannelMember.is_active.is_(True),
)
.first()
).limit(1)
)
membership = result.scalars().first()
if membership:
return self._to_channel_model(channel, db=db)
return await self._to_channel_model(channel, db=db)
else:
return None
# For channels that are NOT group/dm, fall back to ACL-based read access
query = db.query(Channel).filter(Channel.id == id)
stmt = select(Channel).filter(Channel.id == id)
# Determine user groups
user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id, db=db)]
user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id, db=db)]
# Apply ACL rules
query = self._has_permission(
stmt = self._has_permission(
db,
query,
stmt,
{'user_id': user_id, 'group_ids': user_group_ids},
permission='read',
)
channel_allowed = query.first()
return self._to_channel_model(channel_allowed, db=db) if channel_allowed else None
result = await db.execute(stmt)
channel_allowed = result.scalars().first()
return await self._to_channel_model(channel_allowed, db=db) if channel_allowed else None
def update_channel_by_id(
self, id: str, form_data: ChannelForm, db: Optional[Session] = None
async def update_channel_by_id(
self, id: str, form_data: ChannelForm, db: Optional[AsyncSession] = None
) -> Optional[ChannelModel]:
with get_db_context(db) as db:
channel = db.query(Channel).filter(Channel.id == id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(Channel).filter(Channel.id == id))
channel = result.scalars().first()
if not channel:
return None
@ -823,16 +821,16 @@ class ChannelTable:
channel.meta = form_data.meta
if form_data.access_grants is not None:
AccessGrants.set_access_grants('channel', id, form_data.access_grants, db=db)
await AccessGrants.set_access_grants('channel', id, form_data.access_grants, db=db)
channel.updated_at = int(time.time_ns())
db.commit()
return self._to_channel_model(channel, db=db) if channel else None
await db.commit()
return await self._to_channel_model(channel, db=db) if channel else None
def add_file_to_channel_by_id(
self, channel_id: str, file_id: str, user_id: str, db: Optional[Session] = None
async def add_file_to_channel_by_id(
self, channel_id: str, file_id: str, user_id: str, db: Optional[AsyncSession] = None
) -> Optional[ChannelFileModel]:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
channel_file = ChannelFileModel(
**{
'id': str(uuid.uuid4()),
@ -847,8 +845,8 @@ class ChannelTable:
try:
result = ChannelFile(**channel_file.model_dump())
db.add(result)
db.commit()
db.refresh(result)
await db.commit()
await db.refresh(result)
if result:
return ChannelFileModel.model_validate(result)
else:
@ -856,55 +854,58 @@ class ChannelTable:
except Exception:
return None
def set_file_message_id_in_channel_by_id(
async def set_file_message_id_in_channel_by_id(
self,
channel_id: str,
file_id: str,
message_id: str,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> bool:
try:
with get_db_context(db) as db:
channel_file = db.query(ChannelFile).filter_by(channel_id=channel_id, file_id=file_id).first()
async with get_async_db_context(db) as db:
result = await db.execute(
select(ChannelFile).filter_by(channel_id=channel_id, file_id=file_id)
)
channel_file = result.scalars().first()
if not channel_file:
return False
channel_file.message_id = message_id
channel_file.updated_at = int(time.time())
db.commit()
await db.commit()
return True
except Exception:
return False
def remove_file_from_channel_by_id(self, channel_id: str, file_id: str, db: Optional[Session] = None) -> bool:
async def remove_file_from_channel_by_id(self, channel_id: str, file_id: str, db: Optional[AsyncSession] = None) -> bool:
try:
with get_db_context(db) as db:
db.query(ChannelFile).filter_by(channel_id=channel_id, file_id=file_id).delete()
db.commit()
async with get_async_db_context(db) as db:
await db.execute(delete(ChannelFile).filter_by(channel_id=channel_id, file_id=file_id))
await db.commit()
return True
except Exception:
return False
def delete_channel_by_id(self, id: str, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
AccessGrants.revoke_all_access('channel', id, db=db)
db.query(Channel).filter(Channel.id == id).delete()
db.commit()
async def delete_channel_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
await AccessGrants.revoke_all_access('channel', id, db=db)
await db.execute(delete(Channel).filter(Channel.id == id))
await db.commit()
return True
####################
# Webhook Methods
####################
def insert_webhook(
async def insert_webhook(
self,
channel_id: str,
user_id: str,
form_data: ChannelWebhookForm,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[ChannelWebhookModel]:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
webhook = ChannelWebhookModel(
id=str(uuid.uuid4()),
channel_id=channel_id,
@ -917,63 +918,66 @@ class ChannelTable:
updated_at=int(time.time_ns()),
)
db.add(ChannelWebhook(**webhook.model_dump()))
db.commit()
await db.commit()
return webhook
def get_webhooks_by_channel_id(self, channel_id: str, db: Optional[Session] = None) -> list[ChannelWebhookModel]:
with get_db_context(db) as db:
webhooks = db.query(ChannelWebhook).filter(ChannelWebhook.channel_id == channel_id).all()
async def get_webhooks_by_channel_id(self, channel_id: str, db: Optional[AsyncSession] = None) -> list[ChannelWebhookModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(ChannelWebhook).filter(ChannelWebhook.channel_id == channel_id))
webhooks = result.scalars().all()
return [ChannelWebhookModel.model_validate(w) for w in webhooks]
def get_webhook_by_id(self, webhook_id: str, db: Optional[Session] = None) -> Optional[ChannelWebhookModel]:
with get_db_context(db) as db:
webhook = db.query(ChannelWebhook).filter(ChannelWebhook.id == webhook_id).first()
async def get_webhook_by_id(self, webhook_id: str, db: Optional[AsyncSession] = None) -> Optional[ChannelWebhookModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(ChannelWebhook).filter(ChannelWebhook.id == webhook_id))
webhook = result.scalars().first()
return ChannelWebhookModel.model_validate(webhook) if webhook else None
def get_webhook_by_id_and_token(
self, webhook_id: str, token: str, db: Optional[Session] = None
async def get_webhook_by_id_and_token(
self, webhook_id: str, token: str, db: Optional[AsyncSession] = None
) -> Optional[ChannelWebhookModel]:
with get_db_context(db) as db:
webhook = (
db.query(ChannelWebhook)
.filter(
async with get_async_db_context(db) as db:
result = await db.execute(
select(ChannelWebhook).filter(
ChannelWebhook.id == webhook_id,
ChannelWebhook.token == token,
)
.first()
)
webhook = result.scalars().first()
return ChannelWebhookModel.model_validate(webhook) if webhook else None
def update_webhook_by_id(
async def update_webhook_by_id(
self,
webhook_id: str,
form_data: ChannelWebhookForm,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[ChannelWebhookModel]:
with get_db_context(db) as db:
webhook = db.query(ChannelWebhook).filter(ChannelWebhook.id == webhook_id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(ChannelWebhook).filter(ChannelWebhook.id == webhook_id))
webhook = result.scalars().first()
if not webhook:
return None
webhook.name = form_data.name
webhook.profile_image_url = form_data.profile_image_url
webhook.updated_at = int(time.time_ns())
db.commit()
await db.commit()
return ChannelWebhookModel.model_validate(webhook)
def update_webhook_last_used_at(self, webhook_id: str, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
webhook = db.query(ChannelWebhook).filter(ChannelWebhook.id == webhook_id).first()
async def update_webhook_last_used_at(self, webhook_id: str, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
result = await db.execute(select(ChannelWebhook).filter(ChannelWebhook.id == webhook_id))
webhook = result.scalars().first()
if not webhook:
return False
webhook.last_used_at = int(time.time_ns())
db.commit()
await db.commit()
return True
def delete_webhook_by_id(self, webhook_id: str, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
result = db.query(ChannelWebhook).filter(ChannelWebhook.id == webhook_id).delete()
db.commit()
return result > 0
async def delete_webhook_by_id(self, webhook_id: str, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
result = await db.execute(delete(ChannelWebhook).filter(ChannelWebhook.id == webhook_id))
await db.commit()
return result.rowcount > 0
Channels = ChannelTable()

View file

@ -3,8 +3,9 @@ import time
import uuid
from typing import Any, Optional
from sqlalchemy.orm import Session
from open_webui.internal.db import Base, get_db_context
from sqlalchemy import select, delete, func, cast, Integer
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import Base, get_async_db_context
from open_webui.utils.response import normalize_usage
from pydantic import BaseModel, ConfigDict
@ -16,7 +17,6 @@ from sqlalchemy import (
Text,
JSON,
Index,
func,
)
####################
@ -129,23 +129,23 @@ class ChatMessageModel(BaseModel):
class ChatMessageTable:
def upsert_message(
async def upsert_message(
self,
message_id: str,
chat_id: str,
user_id: str,
data: dict,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[ChatMessageModel]:
"""Insert or update a chat message."""
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
now = int(time.time())
timestamp = data.get('timestamp', now)
# Use composite ID: {chat_id}-{message_id}
composite_id = f'{chat_id}-{message_id}'
existing = db.get(ChatMessage, composite_id)
existing = await db.get(ChatMessage, composite_id)
if existing:
# Update existing
if 'role' in data:
@ -178,8 +178,8 @@ class ChatMessageTable:
# from accidentally clearing the primary response's token counts
existing.usage = {**(existing.usage or {}), **usage}
existing.updated_at = now
db.commit()
db.refresh(existing)
await db.commit()
await db.refresh(existing)
return ChatMessageModel.model_validate(existing)
else:
# Insert new
@ -205,143 +205,155 @@ class ChatMessageTable:
updated_at=now,
)
db.add(message)
db.commit()
db.refresh(message)
await db.commit()
await db.refresh(message)
return ChatMessageModel.model_validate(message)
def get_message_by_id(self, id: str, db: Optional[Session] = None) -> Optional[ChatMessageModel]:
with get_db_context(db) as db:
message = db.get(ChatMessage, id)
async def get_message_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[ChatMessageModel]:
async with get_async_db_context(db) as db:
message = await db.get(ChatMessage, id)
return ChatMessageModel.model_validate(message) if message else None
def get_messages_by_chat_id(self, chat_id: str, db: Optional[Session] = None) -> list[ChatMessageModel]:
with get_db_context(db) as db:
messages = db.query(ChatMessage).filter_by(chat_id=chat_id).order_by(ChatMessage.created_at.asc()).all()
async def get_messages_by_chat_id(self, chat_id: str, db: Optional[AsyncSession] = None) -> list[ChatMessageModel]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(ChatMessage).filter_by(chat_id=chat_id).order_by(ChatMessage.created_at.asc())
)
messages = result.scalars().all()
return [ChatMessageModel.model_validate(message) for message in messages]
def get_messages_by_user_id(
async def get_messages_by_user_id(
self,
user_id: str,
skip: int = 0,
limit: int = 50,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> list[ChatMessageModel]:
with get_db_context(db) as db:
messages = (
db.query(ChatMessage)
async with get_async_db_context(db) as db:
result = await db.execute(
select(ChatMessage)
.filter_by(user_id=user_id)
.order_by(ChatMessage.created_at.desc())
.offset(skip)
.limit(limit)
.all()
)
messages = result.scalars().all()
return [ChatMessageModel.model_validate(message) for message in messages]
def get_messages_by_model_id(
async def get_messages_by_model_id(
self,
model_id: str,
start_date: Optional[int] = None,
end_date: Optional[int] = None,
skip: int = 0,
limit: int = 100,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> list[ChatMessageModel]:
with get_db_context(db) as db:
query = db.query(ChatMessage).filter_by(model_id=model_id)
async with get_async_db_context(db) as db:
stmt = select(ChatMessage).filter_by(model_id=model_id)
if start_date:
query = query.filter(ChatMessage.created_at >= start_date)
stmt = stmt.filter(ChatMessage.created_at >= start_date)
if end_date:
query = query.filter(ChatMessage.created_at <= end_date)
messages = query.order_by(ChatMessage.created_at.desc()).offset(skip).limit(limit).all()
stmt = stmt.filter(ChatMessage.created_at <= end_date)
stmt = stmt.order_by(ChatMessage.created_at.desc()).offset(skip).limit(limit)
result = await db.execute(stmt)
messages = result.scalars().all()
return [ChatMessageModel.model_validate(message) for message in messages]
def get_chat_ids_by_model_id(
async def get_chat_ids_by_model_id(
self,
model_id: str,
start_date: Optional[int] = None,
end_date: Optional[int] = None,
skip: int = 0,
limit: int = 50,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> list[str]:
"""Get distinct chat_ids that used a specific model."""
with get_db_context(db) as db:
query = db.query(
ChatMessage.chat_id,
func.max(ChatMessage.created_at).label('last_message_at'),
).filter(ChatMessage.model_id == model_id)
async with get_async_db_context(db) as db:
stmt = (
select(
ChatMessage.chat_id,
func.max(ChatMessage.created_at).label('last_message_at'),
)
.filter(ChatMessage.model_id == model_id)
)
if start_date:
query = query.filter(ChatMessage.created_at >= start_date)
stmt = stmt.filter(ChatMessage.created_at >= start_date)
if end_date:
query = query.filter(ChatMessage.created_at <= end_date)
stmt = stmt.filter(ChatMessage.created_at <= end_date)
# Group by chat_id and order by most recent message in each chat
# Secondary sort on chat_id ensures deterministic pagination
# (prevents duplicates across pages when timestamps tie)
chat_ids = (
query.group_by(ChatMessage.chat_id)
stmt = (
stmt.group_by(ChatMessage.chat_id)
.order_by(func.max(ChatMessage.created_at).desc(), ChatMessage.chat_id)
.offset(skip)
.limit(limit)
.all()
)
result = await db.execute(stmt)
chat_ids = result.all()
return [chat_id for chat_id, _ in chat_ids]
def delete_messages_by_chat_id(self, chat_id: str, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
db.query(ChatMessage).filter_by(chat_id=chat_id).delete()
db.commit()
async def delete_messages_by_chat_id(self, chat_id: str, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
await db.execute(delete(ChatMessage).filter_by(chat_id=chat_id))
await db.commit()
return True
# Analytics methods
def get_message_count_by_model(
async def get_message_count_by_model(
self,
start_date: Optional[int] = None,
end_date: Optional[int] = None,
group_id: Optional[str] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> dict[str, int]:
with get_db_context(db) as db:
from sqlalchemy import func
async with get_async_db_context(db) as db:
from open_webui.models.groups import GroupMember
query = db.query(ChatMessage.model_id, func.count(ChatMessage.id).label('count')).filter(
ChatMessage.role == 'assistant',
ChatMessage.model_id.isnot(None),
~ChatMessage.user_id.like('shared-%'),
stmt = (
select(ChatMessage.model_id, func.count(ChatMessage.id).label('count'))
.filter(
ChatMessage.role == 'assistant',
ChatMessage.model_id.isnot(None),
~ChatMessage.user_id.like('shared-%'),
)
)
if start_date:
query = query.filter(ChatMessage.created_at >= start_date)
stmt = stmt.filter(ChatMessage.created_at >= start_date)
if end_date:
query = query.filter(ChatMessage.created_at <= end_date)
stmt = stmt.filter(ChatMessage.created_at <= end_date)
if group_id:
group_users = db.query(GroupMember.user_id).filter(GroupMember.group_id == group_id).subquery()
query = query.filter(ChatMessage.user_id.in_(group_users))
group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery()
stmt = stmt.filter(ChatMessage.user_id.in_(group_users))
results = query.group_by(ChatMessage.model_id).all()
return {row.model_id: row.count for row in results}
stmt = stmt.group_by(ChatMessage.model_id)
result = await db.execute(stmt)
return {row.model_id: row.count for row in result.all()}
def get_token_usage_by_model(
async def get_token_usage_by_model(
self,
start_date: Optional[int] = None,
end_date: Optional[int] = None,
group_id: Optional[str] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> dict[str, dict]:
"""Aggregate token usage by model using database-level aggregation."""
with get_db_context(db) as db:
from sqlalchemy import func, cast, Integer
async with get_async_db_context(db) as db:
from open_webui.models.groups import GroupMember
dialect = db.bind.dialect.name
# We need the dialect to determine JSON extraction syntax
# For async sessions, access via get_bind()
bind = await db.connection()
dialect = bind.dialect.name
if dialect == 'sqlite':
input_tokens = cast(func.json_extract(ChatMessage.usage, '$.input_tokens'), Integer)
output_tokens = cast(func.json_extract(ChatMessage.usage, '$.output_tokens'), Integer)
elif dialect == 'postgresql':
# Use json_extract_path_text for PostgreSQL JSON columns
input_tokens = cast(
func.json_extract_path_text(ChatMessage.usage, 'input_tokens'),
Integer,
@ -353,27 +365,31 @@ class ChatMessageTable:
else:
raise NotImplementedError(f'Unsupported dialect: {dialect}')
query = db.query(
ChatMessage.model_id,
func.coalesce(func.sum(input_tokens), 0).label('input_tokens'),
func.coalesce(func.sum(output_tokens), 0).label('output_tokens'),
func.count(ChatMessage.id).label('message_count'),
).filter(
ChatMessage.role == 'assistant',
ChatMessage.model_id.isnot(None),
ChatMessage.usage.isnot(None),
~ChatMessage.user_id.like('shared-%'),
stmt = (
select(
ChatMessage.model_id,
func.coalesce(func.sum(input_tokens), 0).label('input_tokens'),
func.coalesce(func.sum(output_tokens), 0).label('output_tokens'),
func.count(ChatMessage.id).label('message_count'),
)
.filter(
ChatMessage.role == 'assistant',
ChatMessage.model_id.isnot(None),
ChatMessage.usage.isnot(None),
~ChatMessage.user_id.like('shared-%'),
)
)
if start_date:
query = query.filter(ChatMessage.created_at >= start_date)
stmt = stmt.filter(ChatMessage.created_at >= start_date)
if end_date:
query = query.filter(ChatMessage.created_at <= end_date)
stmt = stmt.filter(ChatMessage.created_at <= end_date)
if group_id:
group_users = db.query(GroupMember.user_id).filter(GroupMember.group_id == group_id).subquery()
query = query.filter(ChatMessage.user_id.in_(group_users))
group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery()
stmt = stmt.filter(ChatMessage.user_id.in_(group_users))
results = query.group_by(ChatMessage.model_id).all()
stmt = stmt.group_by(ChatMessage.model_id)
result = await db.execute(stmt)
return {
row.model_id: {
@ -382,28 +398,27 @@ class ChatMessageTable:
'total_tokens': row.input_tokens + row.output_tokens,
'message_count': row.message_count,
}
for row in results
for row in result.all()
}
def get_token_usage_by_user(
async def get_token_usage_by_user(
self,
start_date: Optional[int] = None,
end_date: Optional[int] = None,
group_id: Optional[str] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> dict[str, dict]:
"""Aggregate token usage by user using database-level aggregation."""
with get_db_context(db) as db:
from sqlalchemy import func, cast, Integer
async with get_async_db_context(db) as db:
from open_webui.models.groups import GroupMember
dialect = db.bind.dialect.name
bind = await db.connection()
dialect = bind.dialect.name
if dialect == 'sqlite':
input_tokens = cast(func.json_extract(ChatMessage.usage, '$.input_tokens'), Integer)
output_tokens = cast(func.json_extract(ChatMessage.usage, '$.output_tokens'), Integer)
elif dialect == 'postgresql':
# Use json_extract_path_text for PostgreSQL JSON columns
input_tokens = cast(
func.json_extract_path_text(ChatMessage.usage, 'input_tokens'),
Integer,
@ -415,27 +430,31 @@ class ChatMessageTable:
else:
raise NotImplementedError(f'Unsupported dialect: {dialect}')
query = db.query(
ChatMessage.user_id,
func.coalesce(func.sum(input_tokens), 0).label('input_tokens'),
func.coalesce(func.sum(output_tokens), 0).label('output_tokens'),
func.count(ChatMessage.id).label('message_count'),
).filter(
ChatMessage.role == 'assistant',
ChatMessage.user_id.isnot(None),
ChatMessage.usage.isnot(None),
~ChatMessage.user_id.like('shared-%'),
stmt = (
select(
ChatMessage.user_id,
func.coalesce(func.sum(input_tokens), 0).label('input_tokens'),
func.coalesce(func.sum(output_tokens), 0).label('output_tokens'),
func.count(ChatMessage.id).label('message_count'),
)
.filter(
ChatMessage.role == 'assistant',
ChatMessage.user_id.isnot(None),
ChatMessage.usage.isnot(None),
~ChatMessage.user_id.like('shared-%'),
)
)
if start_date:
query = query.filter(ChatMessage.created_at >= start_date)
stmt = stmt.filter(ChatMessage.created_at >= start_date)
if end_date:
query = query.filter(ChatMessage.created_at <= end_date)
stmt = stmt.filter(ChatMessage.created_at <= end_date)
if group_id:
group_users = db.query(GroupMember.user_id).filter(GroupMember.group_id == group_id).subquery()
query = query.filter(ChatMessage.user_id.in_(group_users))
group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery()
stmt = stmt.filter(ChatMessage.user_id.in_(group_users))
results = query.group_by(ChatMessage.user_id).all()
stmt = stmt.group_by(ChatMessage.user_id)
result = await db.execute(stmt)
return {
row.user_id: {
@ -444,88 +463,94 @@ class ChatMessageTable:
'total_tokens': row.input_tokens + row.output_tokens,
'message_count': row.message_count,
}
for row in results
for row in result.all()
}
def get_message_count_by_user(
async def get_message_count_by_user(
self,
start_date: Optional[int] = None,
end_date: Optional[int] = None,
group_id: Optional[str] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> dict[str, int]:
with get_db_context(db) as db:
from sqlalchemy import func
async with get_async_db_context(db) as db:
from open_webui.models.groups import GroupMember
query = db.query(ChatMessage.user_id, func.count(ChatMessage.id).label('count')).filter(
~ChatMessage.user_id.like('shared-%')
stmt = (
select(ChatMessage.user_id, func.count(ChatMessage.id).label('count'))
.filter(~ChatMessage.user_id.like('shared-%'))
)
if start_date:
query = query.filter(ChatMessage.created_at >= start_date)
stmt = stmt.filter(ChatMessage.created_at >= start_date)
if end_date:
query = query.filter(ChatMessage.created_at <= end_date)
stmt = stmt.filter(ChatMessage.created_at <= end_date)
if group_id:
group_users = db.query(GroupMember.user_id).filter(GroupMember.group_id == group_id).subquery()
query = query.filter(ChatMessage.user_id.in_(group_users))
group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery()
stmt = stmt.filter(ChatMessage.user_id.in_(group_users))
results = query.group_by(ChatMessage.user_id).all()
return {row.user_id: row.count for row in results}
stmt = stmt.group_by(ChatMessage.user_id)
result = await db.execute(stmt)
return {row.user_id: row.count for row in result.all()}
def get_message_count_by_chat(
async def get_message_count_by_chat(
self,
start_date: Optional[int] = None,
end_date: Optional[int] = None,
group_id: Optional[str] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> dict[str, int]:
with get_db_context(db) as db:
from sqlalchemy import func
async with get_async_db_context(db) as db:
from open_webui.models.groups import GroupMember
query = db.query(ChatMessage.chat_id, func.count(ChatMessage.id).label('count')).filter(
~ChatMessage.user_id.like('shared-%')
stmt = (
select(ChatMessage.chat_id, func.count(ChatMessage.id).label('count'))
.filter(~ChatMessage.user_id.like('shared-%'))
)
if start_date:
query = query.filter(ChatMessage.created_at >= start_date)
stmt = stmt.filter(ChatMessage.created_at >= start_date)
if end_date:
query = query.filter(ChatMessage.created_at <= end_date)
stmt = stmt.filter(ChatMessage.created_at <= end_date)
if group_id:
group_users = db.query(GroupMember.user_id).filter(GroupMember.group_id == group_id).subquery()
query = query.filter(ChatMessage.user_id.in_(group_users))
group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery()
stmt = stmt.filter(ChatMessage.user_id.in_(group_users))
results = query.group_by(ChatMessage.chat_id).all()
return {row.chat_id: row.count for row in results}
stmt = stmt.group_by(ChatMessage.chat_id)
result = await db.execute(stmt)
return {row.chat_id: row.count for row in result.all()}
def get_daily_message_counts_by_model(
async def get_daily_message_counts_by_model(
self,
start_date: Optional[int] = None,
end_date: Optional[int] = None,
group_id: Optional[str] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> dict[str, dict[str, int]]:
"""Get message counts grouped by day and model."""
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
from datetime import datetime, timedelta
from open_webui.models.groups import GroupMember
query = db.query(ChatMessage.created_at, ChatMessage.model_id).filter(
ChatMessage.role == 'assistant',
ChatMessage.model_id.isnot(None),
~ChatMessage.user_id.like('shared-%'),
stmt = (
select(ChatMessage.created_at, ChatMessage.model_id)
.filter(
ChatMessage.role == 'assistant',
ChatMessage.model_id.isnot(None),
~ChatMessage.user_id.like('shared-%'),
)
)
if start_date:
query = query.filter(ChatMessage.created_at >= start_date)
stmt = stmt.filter(ChatMessage.created_at >= start_date)
if end_date:
query = query.filter(ChatMessage.created_at <= end_date)
stmt = stmt.filter(ChatMessage.created_at <= end_date)
if group_id:
group_users = db.query(GroupMember.user_id).filter(GroupMember.group_id == group_id).subquery()
query = query.filter(ChatMessage.user_id.in_(group_users))
group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery()
stmt = stmt.filter(ChatMessage.user_id.in_(group_users))
results = query.all()
result = await db.execute(stmt)
results = result.all()
# Group by date -> model -> count
daily_counts: dict[str, dict[str, int]] = {}
@ -547,28 +572,32 @@ class ChatMessageTable:
return daily_counts
def get_hourly_message_counts_by_model(
async def get_hourly_message_counts_by_model(
self,
start_date: Optional[int] = None,
end_date: Optional[int] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> dict[str, dict[str, int]]:
"""Get message counts grouped by hour and model."""
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
from datetime import datetime, timedelta
query = db.query(ChatMessage.created_at, ChatMessage.model_id).filter(
ChatMessage.role == 'assistant',
ChatMessage.model_id.isnot(None),
~ChatMessage.user_id.like('shared-%'),
stmt = (
select(ChatMessage.created_at, ChatMessage.model_id)
.filter(
ChatMessage.role == 'assistant',
ChatMessage.model_id.isnot(None),
~ChatMessage.user_id.like('shared-%'),
)
)
if start_date:
query = query.filter(ChatMessage.created_at >= start_date)
stmt = stmt.filter(ChatMessage.created_at >= start_date)
if end_date:
query = query.filter(ChatMessage.created_at <= end_date)
stmt = stmt.filter(ChatMessage.created_at <= end_date)
results = query.all()
result = await db.execute(stmt)
results = result.all()
# Group by hour -> model -> count
hourly_counts: dict[str, dict[str, int]] = {}

File diff suppressed because it is too large Load diff

View file

@ -3,9 +3,10 @@ import time
import uuid
from typing import Optional
from sqlalchemy.orm import Session
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
from open_webui.models.users import User
from sqlalchemy import select, delete, func
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import Base, JSONField, get_async_db_context
from open_webui.models.users import User, UserModel
from pydantic import BaseModel, ConfigDict
from sqlalchemy import BigInteger, Column, Text, JSON, Boolean
@ -139,10 +140,10 @@ class ModelHistoryResponse(BaseModel):
class FeedbackTable:
def insert_new_feedback(
self, user_id: str, form_data: FeedbackForm, db: Optional[Session] = None
async def insert_new_feedback(
self, user_id: str, form_data: FeedbackForm, db: Optional[AsyncSession] = None
) -> Optional[FeedbackModel]:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
id = str(uuid.uuid4())
feedback = FeedbackModel(
**{
@ -157,8 +158,8 @@ class FeedbackTable:
try:
result = Feedback(**feedback.model_dump())
db.add(result)
db.commit()
db.refresh(result)
await db.commit()
await db.refresh(result)
if result:
return FeedbackModel.model_validate(result)
else:
@ -167,97 +168,101 @@ class FeedbackTable:
log.exception(f'Error creating a new feedback: {e}')
return None
def get_feedback_by_id(self, id: str, db: Optional[Session] = None) -> Optional[FeedbackModel]:
async def get_feedback_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[FeedbackModel]:
try:
with get_db_context(db) as db:
feedback = db.query(Feedback).filter_by(id=id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(Feedback).filter_by(id=id))
feedback = result.scalars().first()
if not feedback:
return None
return FeedbackModel.model_validate(feedback)
except Exception:
return None
def get_feedback_by_id_and_user_id(
self, id: str, user_id: str, db: Optional[Session] = None
async def get_feedback_by_id_and_user_id(
self, id: str, user_id: str, db: Optional[AsyncSession] = None
) -> Optional[FeedbackModel]:
try:
with get_db_context(db) as db:
feedback = db.query(Feedback).filter_by(id=id, user_id=user_id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(Feedback).filter_by(id=id, user_id=user_id))
feedback = result.scalars().first()
if not feedback:
return None
return FeedbackModel.model_validate(feedback)
except Exception:
return None
def get_feedbacks_by_chat_id(self, chat_id: str, db: Optional[Session] = None) -> list[FeedbackModel]:
async def get_feedbacks_by_chat_id(self, chat_id: str, db: Optional[AsyncSession] = None) -> list[FeedbackModel]:
"""Get all feedbacks for a specific chat."""
try:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
# meta.chat_id stores the chat reference
feedbacks = (
db.query(Feedback)
result = await db.execute(
select(Feedback)
.filter(Feedback.meta['chat_id'].as_string() == chat_id)
.order_by(Feedback.created_at.desc())
.all()
)
feedbacks = result.scalars().all()
return [FeedbackModel.model_validate(fb) for fb in feedbacks]
except Exception:
return []
def get_feedback_items(
async def get_feedback_items(
self,
filter: dict = {},
skip: int = 0,
limit: int = 30,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> FeedbackListResponse:
with get_db_context(db) as db:
query = db.query(Feedback, User).join(User, Feedback.user_id == User.id)
async with get_async_db_context(db) as db:
stmt = select(Feedback, User).join(User, Feedback.user_id == User.id)
if filter:
# Apply model_id filter (exact match)
model_id = filter.get('model_id')
if model_id:
query = query.filter(Feedback.data['model_id'].as_string() == model_id)
stmt = stmt.filter(Feedback.data['model_id'].as_string() == model_id)
order_by = filter.get('order_by')
direction = filter.get('direction')
if order_by == 'username':
if direction == 'asc':
query = query.order_by(User.name.asc())
stmt = stmt.order_by(User.name.asc())
else:
query = query.order_by(User.name.desc())
stmt = stmt.order_by(User.name.desc())
elif order_by == 'model_id':
# it's stored in feedback.data['model_id']
if direction == 'asc':
query = query.order_by(Feedback.data['model_id'].as_string().asc())
stmt = stmt.order_by(Feedback.data['model_id'].as_string().asc())
else:
query = query.order_by(Feedback.data['model_id'].as_string().desc())
stmt = stmt.order_by(Feedback.data['model_id'].as_string().desc())
elif order_by == 'rating':
# it's stored in feedback.data['rating']
if direction == 'asc':
query = query.order_by(Feedback.data['rating'].as_string().asc())
stmt = stmt.order_by(Feedback.data['rating'].as_string().asc())
else:
query = query.order_by(Feedback.data['rating'].as_string().desc())
stmt = stmt.order_by(Feedback.data['rating'].as_string().desc())
elif order_by == 'updated_at':
if direction == 'asc':
query = query.order_by(Feedback.updated_at.asc())
stmt = stmt.order_by(Feedback.updated_at.asc())
else:
query = query.order_by(Feedback.updated_at.desc())
stmt = stmt.order_by(Feedback.updated_at.desc())
else:
query = query.order_by(Feedback.created_at.desc())
stmt = stmt.order_by(Feedback.created_at.desc())
# Count BEFORE pagination
total = query.count()
count_result = await db.execute(
select(func.count()).select_from(stmt.subquery())
)
total = count_result.scalar()
if skip:
query = query.offset(skip)
stmt = stmt.offset(skip)
if limit:
query = query.limit(limit)
stmt = stmt.limit(limit)
items = query.all()
result = await db.execute(stmt)
items = result.all()
feedbacks = []
for feedback, user in items:
@ -267,15 +272,17 @@ class FeedbackTable:
return FeedbackListResponse(items=feedbacks, total=total)
def get_all_feedbacks(self, db: Optional[Session] = None) -> list[FeedbackModel]:
with get_db_context(db) as db:
return [
FeedbackModel.model_validate(feedback)
for feedback in db.query(Feedback).order_by(Feedback.updated_at.desc()).all()
]
async def get_all_feedbacks(self, db: Optional[AsyncSession] = None) -> list[FeedbackModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(Feedback).order_by(Feedback.updated_at.desc()))
return [FeedbackModel.model_validate(feedback) for feedback in result.scalars().all()]
def get_all_feedback_ids(self, db: Optional[Session] = None) -> list[FeedbackIdResponse]:
with get_db_context(db) as db:
async def get_all_feedback_ids(self, db: Optional[AsyncSession] = None) -> list[FeedbackIdResponse]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(Feedback.id, Feedback.user_id, Feedback.created_at, Feedback.updated_at)
.order_by(Feedback.updated_at.desc())
)
return [
FeedbackIdResponse(
id=row.id,
@ -283,36 +290,28 @@ class FeedbackTable:
created_at=row.created_at,
updated_at=row.updated_at,
)
for row in db.query(
Feedback.id,
Feedback.user_id,
Feedback.created_at,
Feedback.updated_at,
)
.order_by(Feedback.updated_at.desc())
.all()
for row in result.all()
]
def get_distinct_model_ids(self, db: Optional[Session] = None) -> list[str]:
async def get_distinct_model_ids(self, db: Optional[AsyncSession] = None) -> list[str]:
"""Get distinct model_ids from feedback data for filter dropdowns."""
with get_db_context(db) as db:
rows = (
db.query(Feedback.data['model_id'].as_string())
async with get_async_db_context(db) as db:
result = await db.execute(
select(Feedback.data['model_id'].as_string())
.filter(Feedback.data['model_id'].as_string().isnot(None))
.distinct()
.all()
)
rows = result.all()
return sorted([row[0] for row in rows if row[0]])
def get_feedbacks_for_leaderboard(self, db: Optional[Session] = None) -> list[LeaderboardFeedbackData]:
async def get_feedbacks_for_leaderboard(self, db: Optional[AsyncSession] = None) -> list[LeaderboardFeedbackData]:
"""Fetch only id and data for leaderboard computation (excludes snapshot/meta)."""
with get_db_context(db) as db:
return [
LeaderboardFeedbackData(id=row.id, data=row.data) for row in db.query(Feedback.id, Feedback.data).all()
]
async with get_async_db_context(db) as db:
result = await db.execute(select(Feedback.id, Feedback.data))
return [LeaderboardFeedbackData(id=row.id, data=row.data) for row in result.all()]
def get_model_evaluation_history(
self, model_id: str, days: int = 30, db: Optional[Session] = None
async def get_model_evaluation_history(
self, model_id: str, days: int = 30, db: Optional[AsyncSession] = None
) -> list[ModelHistoryEntry]:
"""
Get daily wins/losses for a specific model over the past N days.
@ -322,13 +321,16 @@ class FeedbackTable:
from datetime import datetime, timedelta
from collections import defaultdict
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
if days == 0:
# All time - no cutoff
rows = db.query(Feedback.created_at, Feedback.data).all()
result = await db.execute(select(Feedback.created_at, Feedback.data))
else:
cutoff = int(time.time()) - (days * 86400)
rows = db.query(Feedback.created_at, Feedback.data).filter(Feedback.created_at >= cutoff).all()
result = await db.execute(
select(Feedback.created_at, Feedback.data).filter(Feedback.created_at >= cutoff)
)
rows = result.all()
daily_counts = defaultdict(lambda: {'won': 0, 'lost': 0})
first_date = None
@ -374,25 +376,26 @@ class FeedbackTable:
return result
def get_feedbacks_by_type(self, type: str, db: Optional[Session] = None) -> list[FeedbackModel]:
with get_db_context(db) as db:
return [
FeedbackModel.model_validate(feedback)
for feedback in db.query(Feedback).filter_by(type=type).order_by(Feedback.updated_at.desc()).all()
]
async def get_feedbacks_by_type(self, type: str, db: Optional[AsyncSession] = None) -> list[FeedbackModel]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(Feedback).filter_by(type=type).order_by(Feedback.updated_at.desc())
)
return [FeedbackModel.model_validate(feedback) for feedback in result.scalars().all()]
def get_feedbacks_by_user_id(self, user_id: str, db: Optional[Session] = None) -> list[FeedbackModel]:
with get_db_context(db) as db:
return [
FeedbackModel.model_validate(feedback)
for feedback in db.query(Feedback).filter_by(user_id=user_id).order_by(Feedback.updated_at.desc()).all()
]
async def get_feedbacks_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[FeedbackModel]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(Feedback).filter_by(user_id=user_id).order_by(Feedback.updated_at.desc())
)
return [FeedbackModel.model_validate(feedback) for feedback in result.scalars().all()]
def update_feedback_by_id(
self, id: str, form_data: FeedbackForm, db: Optional[Session] = None
async def update_feedback_by_id(
self, id: str, form_data: FeedbackForm, db: Optional[AsyncSession] = None
) -> Optional[FeedbackModel]:
with get_db_context(db) as db:
feedback = db.query(Feedback).filter_by(id=id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(Feedback).filter_by(id=id))
feedback = result.scalars().first()
if not feedback:
return None
@ -405,18 +408,19 @@ class FeedbackTable:
feedback.updated_at = int(time.time())
db.commit()
await db.commit()
return FeedbackModel.model_validate(feedback)
def update_feedback_by_id_and_user_id(
async def update_feedback_by_id_and_user_id(
self,
id: str,
user_id: str,
form_data: FeedbackForm,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[FeedbackModel]:
with get_db_context(db) as db:
feedback = db.query(Feedback).filter_by(id=id, user_id=user_id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(Feedback).filter_by(id=id, user_id=user_id))
feedback = result.scalars().first()
if not feedback:
return None
@ -429,38 +433,40 @@ class FeedbackTable:
feedback.updated_at = int(time.time())
db.commit()
await db.commit()
return FeedbackModel.model_validate(feedback)
def delete_feedback_by_id(self, id: str, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
feedback = db.query(Feedback).filter_by(id=id).first()
async def delete_feedback_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
result = await db.execute(select(Feedback).filter_by(id=id))
feedback = result.scalars().first()
if not feedback:
return False
db.delete(feedback)
db.commit()
await db.delete(feedback)
await db.commit()
return True
def delete_feedback_by_id_and_user_id(self, id: str, user_id: str, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
feedback = db.query(Feedback).filter_by(id=id, user_id=user_id).first()
async def delete_feedback_by_id_and_user_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
result = await db.execute(select(Feedback).filter_by(id=id, user_id=user_id))
feedback = result.scalars().first()
if not feedback:
return False
db.delete(feedback)
db.commit()
await db.delete(feedback)
await db.commit()
return True
def delete_feedbacks_by_user_id(self, user_id: str, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
result = db.query(Feedback).filter_by(user_id=user_id).delete()
db.commit()
return result > 0
async def delete_feedbacks_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
result = await db.execute(delete(Feedback).filter_by(user_id=user_id))
await db.commit()
return result.rowcount > 0
def delete_all_feedbacks(self, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
result = db.query(Feedback).delete()
db.commit()
return result > 0
async def delete_all_feedbacks(self, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
result = await db.execute(delete(Feedback))
await db.commit()
return result.rowcount > 0
Feedbacks = FeedbackTable()

View file

@ -2,8 +2,9 @@ import logging
import time
from typing import Optional
from sqlalchemy.orm import Session
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
from sqlalchemy import select, delete, func
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import Base, JSONField, get_async_db_context
from open_webui.utils.misc import sanitize_metadata
from pydantic import BaseModel, ConfigDict, model_validator
from sqlalchemy import BigInteger, Column, String, Text, JSON
@ -124,8 +125,8 @@ class FileUpdateForm(BaseModel):
class FilesTable:
def insert_new_file(self, user_id: str, form_data: FileForm, db: Optional[Session] = None) -> Optional[FileModel]:
with get_db_context(db) as db:
async def insert_new_file(self, user_id: str, form_data: FileForm, db: Optional[AsyncSession] = None) -> Optional[FileModel]:
async with get_async_db_context(db) as db:
file_data = form_data.model_dump()
# Sanitize meta to remove non-JSON-serializable objects
@ -145,8 +146,8 @@ class FilesTable:
try:
result = File(**file.model_dump())
db.add(result)
db.commit()
db.refresh(result)
await db.commit()
await db.refresh(result)
if result:
return FileModel.model_validate(result)
else:
@ -155,21 +156,22 @@ class FilesTable:
log.exception(f'Error inserting a new file: {e}')
return None
def get_file_by_id(self, id: str, db: Optional[Session] = None) -> Optional[FileModel]:
async def get_file_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[FileModel]:
try:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
try:
file = db.get(File, id)
return FileModel.model_validate(file)
file = await db.get(File, id)
return FileModel.model_validate(file) if file else None
except Exception:
return None
except Exception:
return None
def get_file_by_id_and_user_id(self, id: str, user_id: str, db: Optional[Session] = None) -> Optional[FileModel]:
with get_db_context(db) as db:
async def get_file_by_id_and_user_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> Optional[FileModel]:
async with get_async_db_context(db) as db:
try:
file = db.query(File).filter_by(id=id, user_id=user_id).first()
result = await db.execute(select(File).filter_by(id=id, user_id=user_id))
file = result.scalars().first()
if file:
return FileModel.model_validate(file)
else:
@ -177,10 +179,12 @@ class FilesTable:
except Exception:
return None
def get_file_metadata_by_id(self, id: str, db: Optional[Session] = None) -> Optional[FileMetadataResponse]:
with get_db_context(db) as db:
async def get_file_metadata_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[FileMetadataResponse]:
async with get_async_db_context(db) as db:
try:
file = db.get(File, id)
file = await db.get(File, id)
if not file:
return None
return FileMetadataResponse(
id=file.id,
hash=file.hash,
@ -191,12 +195,13 @@ class FilesTable:
except Exception:
return None
def get_files(self, db: Optional[Session] = None) -> list[FileModel]:
with get_db_context(db) as db:
return [FileModel.model_validate(file) for file in db.query(File).all()]
async def get_files(self, db: Optional[AsyncSession] = None) -> list[FileModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(File))
return [FileModel.model_validate(file) for file in result.scalars().all()]
def check_access_by_user_id(self, id, user_id, permission='write', db: Optional[Session] = None) -> bool:
file = self.get_file_by_id(id, db=db)
async def check_access_by_user_id(self, id, user_id, permission='write', db: Optional[AsyncSession] = None) -> bool:
file = await self.get_file_by_id(id, db=db)
if not file:
return False
if file.user_id == user_id:
@ -204,50 +209,59 @@ class FilesTable:
# Implement additional access control logic here as needed
return False
def get_files_by_ids(self, ids: list[str], db: Optional[Session] = None) -> list[FileModel]:
with get_db_context(db) as db:
return [
FileModel.model_validate(file)
for file in db.query(File).filter(File.id.in_(ids)).order_by(File.updated_at.desc()).all()
]
async def get_files_by_ids(self, ids: list[str], db: Optional[AsyncSession] = None) -> list[FileModel]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(File).filter(File.id.in_(ids)).order_by(File.updated_at.desc())
)
return [FileModel.model_validate(file) for file in result.scalars().all()]
def get_file_metadatas_by_ids(self, ids: list[str], db: Optional[Session] = None) -> list[FileMetadataResponse]:
with get_db_context(db) as db:
return [
FileMetadataResponse(
id=file.id,
hash=file.hash,
meta=file.meta,
created_at=file.created_at,
updated_at=file.updated_at,
)
for file in db.query(File.id, File.hash, File.meta, File.created_at, File.updated_at)
async def get_file_metadatas_by_ids(self, ids: list[str], db: Optional[AsyncSession] = None) -> list[FileMetadataResponse]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(File.id, File.hash, File.meta, File.created_at, File.updated_at)
.filter(File.id.in_(ids))
.order_by(File.updated_at.desc())
.all()
)
return [
FileMetadataResponse(
id=row.id,
hash=row.hash,
meta=row.meta,
created_at=row.created_at,
updated_at=row.updated_at,
)
for row in result.all()
]
def get_files_by_user_id(self, user_id: str, db: Optional[Session] = None) -> list[FileModel]:
with get_db_context(db) as db:
return [FileModel.model_validate(file) for file in db.query(File).filter_by(user_id=user_id).all()]
async def get_files_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[FileModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(File).filter_by(user_id=user_id))
return [FileModel.model_validate(file) for file in result.scalars().all()]
def get_file_list(
async def get_file_list(
self,
user_id: Optional[str] = None,
skip: int = 0,
limit: int = 50,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> 'FileListResponse':
with get_db_context(db) as db:
query = db.query(File)
async with get_async_db_context(db) as db:
stmt = select(File)
if user_id:
query = query.filter_by(user_id=user_id)
stmt = stmt.filter_by(user_id=user_id)
total = query.count()
count_result = await db.execute(
select(func.count()).select_from(stmt.subquery())
)
total = count_result.scalar()
result = await db.execute(
stmt.order_by(File.updated_at.desc(), File.id.desc()).offset(skip).limit(limit)
)
items = [
FileModelResponse.model_validate(file, from_attributes=True)
for file in query.order_by(File.updated_at.desc(), File.id.desc()).offset(skip).limit(limit).all()
for file in result.scalars().all()
]
return FileListResponse(items=items, total=total)
@ -275,13 +289,13 @@ class FilesTable:
pattern = pattern.replace('?', '_')
return pattern
def search_files(
async def search_files(
self,
user_id: Optional[str] = None,
filename: str = '*',
skip: int = 0,
limit: int = 100,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> list[FileModel]:
"""
Search files with glob pattern matching, optional user filter, and pagination.
@ -296,27 +310,28 @@ class FilesTable:
Returns:
List of matching FileModel objects, ordered by created_at descending.
"""
with get_db_context(db) as db:
query = db.query(File)
async with get_async_db_context(db) as db:
stmt = select(File)
if user_id:
query = query.filter_by(user_id=user_id)
stmt = stmt.filter_by(user_id=user_id)
pattern = self._glob_to_like_pattern(filename)
if pattern != '%':
query = query.filter(File.filename.ilike(pattern, escape='\\'))
stmt = stmt.filter(File.filename.ilike(pattern, escape='\\'))
return [
FileModel.model_validate(file)
for file in query.order_by(File.created_at.desc(), File.id.desc()).offset(skip).limit(limit).all()
]
result = await db.execute(
stmt.order_by(File.created_at.desc(), File.id.desc()).offset(skip).limit(limit)
)
return [FileModel.model_validate(file) for file in result.scalars().all()]
def update_file_by_id(
self, id: str, form_data: FileUpdateForm, db: Optional[Session] = None
async def update_file_by_id(
self, id: str, form_data: FileUpdateForm, db: Optional[AsyncSession] = None
) -> Optional[FileModel]:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
try:
file = db.query(File).filter_by(id=id).first()
result = await db.execute(select(File).filter_by(id=id))
file = result.scalars().first()
if form_data.hash is not None:
file.hash = form_data.hash
@ -328,63 +343,64 @@ class FilesTable:
file.meta = {**(file.meta if file.meta else {}), **form_data.meta}
file.updated_at = int(time.time())
db.commit()
await db.commit()
return FileModel.model_validate(file)
except Exception as e:
log.exception(f'Error updating file completely by id: {e}')
return None
def update_file_hash_by_id(self, id: str, hash: Optional[str], db: Optional[Session] = None) -> Optional[FileModel]:
with get_db_context(db) as db:
async def update_file_hash_by_id(self, id: str, hash: Optional[str], db: Optional[AsyncSession] = None) -> Optional[FileModel]:
async with get_async_db_context(db) as db:
try:
file = db.query(File).filter_by(id=id).first()
result = await db.execute(select(File).filter_by(id=id))
file = result.scalars().first()
file.hash = hash
file.updated_at = int(time.time())
db.commit()
await db.commit()
return FileModel.model_validate(file)
except Exception:
return None
def update_file_data_by_id(self, id: str, data: dict, db: Optional[Session] = None) -> Optional[FileModel]:
with get_db_context(db) as db:
async def update_file_data_by_id(self, id: str, data: dict, db: Optional[AsyncSession] = None) -> Optional[FileModel]:
async with get_async_db_context(db) as db:
try:
file = db.query(File).filter_by(id=id).first()
result = await db.execute(select(File).filter_by(id=id))
file = result.scalars().first()
file.data = {**(file.data if file.data else {}), **data}
file.updated_at = int(time.time())
db.commit()
await db.commit()
return FileModel.model_validate(file)
except Exception as e:
return None
def update_file_metadata_by_id(self, id: str, meta: dict, db: Optional[Session] = None) -> Optional[FileModel]:
with get_db_context(db) as db:
async def update_file_metadata_by_id(self, id: str, meta: dict, db: Optional[AsyncSession] = None) -> Optional[FileModel]:
async with get_async_db_context(db) as db:
try:
file = db.query(File).filter_by(id=id).first()
result = await db.execute(select(File).filter_by(id=id))
file = result.scalars().first()
file.meta = {**(file.meta if file.meta else {}), **meta}
file.updated_at = int(time.time())
db.commit()
await db.commit()
return FileModel.model_validate(file)
except Exception:
return None
return False
def delete_file_by_id(self, id: str, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
async def delete_file_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
try:
db.query(File).filter_by(id=id).delete()
db.commit()
await db.execute(delete(File).filter_by(id=id))
await db.commit()
return True
except Exception:
return False
def delete_all_files(self, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
async def delete_all_files(self, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
try:
db.query(File).delete()
db.commit()
await db.execute(delete(File))
await db.commit()
return True
except Exception:

View file

@ -6,10 +6,10 @@ import re
from pydantic import BaseModel, ConfigDict
from sqlalchemy import BigInteger, Column, Text, JSON, Boolean, func
from sqlalchemy.orm import Session
from sqlalchemy import BigInteger, Column, Text, JSON, Boolean, func, select, delete
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
from open_webui.internal.db import Base, JSONField, get_async_db_context
log = logging.getLogger(__name__)
@ -85,14 +85,14 @@ class FolderUpdateForm(BaseModel):
class FolderTable:
def insert_new_folder(
async def insert_new_folder(
self,
user_id: str,
form_data: FolderForm,
parent_id: Optional[str] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[FolderModel]:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
id = str(uuid.uuid4())
folder = FolderModel(
**{
@ -107,8 +107,8 @@ class FolderTable:
try:
result = Folder(**folder.model_dump())
db.add(result)
db.commit()
db.refresh(result)
await db.commit()
await db.refresh(result)
if result:
return FolderModel.model_validate(result)
else:
@ -117,12 +117,13 @@ class FolderTable:
log.exception(f'Error inserting a new folder: {e}')
return None
def get_folder_by_id_and_user_id(
self, id: str, user_id: str, db: Optional[Session] = None
async def get_folder_by_id_and_user_id(
self, id: str, user_id: str, db: Optional[AsyncSession] = None
) -> Optional[FolderModel]:
try:
with get_db_context(db) as db:
folder = db.query(Folder).filter_by(id=id, user_id=user_id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(Folder).filter_by(id=id, user_id=user_id))
folder = result.scalars().first()
if not folder:
return None
@ -131,48 +132,50 @@ class FolderTable:
except Exception:
return None
def get_children_folders_by_id_and_user_id(
self, id: str, user_id: str, db: Optional[Session] = None
async def get_children_folders_by_id_and_user_id(
self, id: str, user_id: str, db: Optional[AsyncSession] = None
) -> Optional[list[FolderModel]]:
try:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
folders = []
def get_children(folder):
children = self.get_folders_by_parent_id_and_user_id(folder.id, user_id, db=db)
async def get_children(folder):
children = await self.get_folders_by_parent_id_and_user_id(folder.id, user_id, db=db)
for child in children:
get_children(child)
await get_children(child)
folders.append(child)
folder = db.query(Folder).filter_by(id=id, user_id=user_id).first()
result = await db.execute(select(Folder).filter_by(id=id, user_id=user_id))
folder = result.scalars().first()
if not folder:
return None
get_children(folder)
await get_children(folder)
return folders
except Exception:
return None
def get_folders_by_user_id(self, user_id: str, db: Optional[Session] = None) -> list[FolderModel]:
with get_db_context(db) as db:
return [FolderModel.model_validate(folder) for folder in db.query(Folder).filter_by(user_id=user_id).all()]
async def get_folders_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[FolderModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(Folder).filter_by(user_id=user_id))
return [FolderModel.model_validate(folder) for folder in result.scalars().all()]
def get_folder_by_parent_id_and_user_id_and_name(
async def get_folder_by_parent_id_and_user_id_and_name(
self,
parent_id: Optional[str],
user_id: str,
name: str,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[FolderModel]:
try:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
# Check if folder exists
folder = (
db.query(Folder)
result = await db.execute(
select(Folder)
.filter_by(parent_id=parent_id, user_id=user_id)
.filter(Folder.name.ilike(name))
.first()
)
folder = result.scalars().first()
if not folder:
return None
@ -182,25 +185,24 @@ class FolderTable:
log.error(f'get_folder_by_parent_id_and_user_id_and_name: {e}')
return None
def get_folders_by_parent_id_and_user_id(
self, parent_id: Optional[str], user_id: str, db: Optional[Session] = None
async def get_folders_by_parent_id_and_user_id(
self, parent_id: Optional[str], user_id: str, db: Optional[AsyncSession] = None
) -> list[FolderModel]:
with get_db_context(db) as db:
return [
FolderModel.model_validate(folder)
for folder in db.query(Folder).filter_by(parent_id=parent_id, user_id=user_id).all()
]
async with get_async_db_context(db) as db:
result = await db.execute(select(Folder).filter_by(parent_id=parent_id, user_id=user_id))
return [FolderModel.model_validate(folder) for folder in result.scalars().all()]
def update_folder_parent_id_by_id_and_user_id(
async def update_folder_parent_id_by_id_and_user_id(
self,
id: str,
user_id: str,
parent_id: str,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[FolderModel]:
try:
with get_db_context(db) as db:
folder = db.query(Folder).filter_by(id=id, user_id=user_id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(Folder).filter_by(id=id, user_id=user_id))
folder = result.scalars().first()
if not folder:
return None
@ -208,38 +210,39 @@ class FolderTable:
folder.parent_id = parent_id
folder.updated_at = int(time.time())
db.commit()
await db.commit()
return FolderModel.model_validate(folder)
except Exception as e:
log.error(f'update_folder: {e}')
return
def update_folder_by_id_and_user_id(
async def update_folder_by_id_and_user_id(
self,
id: str,
user_id: str,
form_data: FolderUpdateForm,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[FolderModel]:
try:
with get_db_context(db) as db:
folder = db.query(Folder).filter_by(id=id, user_id=user_id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(Folder).filter_by(id=id, user_id=user_id))
folder = result.scalars().first()
if not folder:
return None
form_data = form_data.model_dump(exclude_unset=True)
existing_folder = (
db.query(Folder)
existing_result = await db.execute(
select(Folder)
.filter_by(
name=form_data.get('name'),
parent_id=folder.parent_id,
user_id=user_id,
)
.first()
)
existing_folder = existing_result.scalars().first()
if existing_folder and existing_folder.id != id:
return None
@ -258,19 +261,20 @@ class FolderTable:
}
folder.updated_at = int(time.time())
db.commit()
await db.commit()
return FolderModel.model_validate(folder)
except Exception as e:
log.error(f'update_folder: {e}')
return
def update_folder_is_expanded_by_id_and_user_id(
self, id: str, user_id: str, is_expanded: bool, db: Optional[Session] = None
async def update_folder_is_expanded_by_id_and_user_id(
self, id: str, user_id: str, is_expanded: bool, db: Optional[AsyncSession] = None
) -> Optional[FolderModel]:
try:
with get_db_context(db) as db:
folder = db.query(Folder).filter_by(id=id, user_id=user_id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(Folder).filter_by(id=id, user_id=user_id))
folder = result.scalars().first()
if not folder:
return None
@ -278,37 +282,39 @@ class FolderTable:
folder.is_expanded = is_expanded
folder.updated_at = int(time.time())
db.commit()
await db.commit()
return FolderModel.model_validate(folder)
except Exception as e:
log.error(f'update_folder: {e}')
return
def delete_folder_by_id_and_user_id(self, id: str, user_id: str, db: Optional[Session] = None) -> list[str]:
async def delete_folder_by_id_and_user_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> list[str]:
try:
folder_ids = []
with get_db_context(db) as db:
folder = db.query(Folder).filter_by(id=id, user_id=user_id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(Folder).filter_by(id=id, user_id=user_id))
folder = result.scalars().first()
if not folder:
return folder_ids
folder_ids.append(folder.id)
# Delete all children folders
def delete_children(folder):
folder_children = self.get_folders_by_parent_id_and_user_id(folder.id, user_id, db=db)
async def delete_children(folder):
folder_children = await self.get_folders_by_parent_id_and_user_id(folder.id, user_id, db=db)
for folder_child in folder_children:
delete_children(folder_child)
await delete_children(folder_child)
folder_ids.append(folder_child.id)
folder = db.query(Folder).filter_by(id=folder_child.id).first()
db.delete(folder)
db.commit()
child_result = await db.execute(select(Folder).filter_by(id=folder_child.id))
child_folder = child_result.scalars().first()
await db.delete(child_folder)
await db.commit()
delete_children(folder)
db.delete(folder)
db.commit()
await delete_children(folder)
await db.delete(folder)
await db.commit()
return folder_ids
except Exception as e:
log.error(f'delete_folder: {e}')
@ -319,8 +325,8 @@ class FolderTable:
name = re.sub(r'[\s_]+', ' ', name)
return name.strip().lower()
def search_folders_by_names(
self, user_id: str, queries: list[str], db: Optional[Session] = None
async def search_folders_by_names(
self, user_id: str, queries: list[str], db: Optional[AsyncSession] = None
) -> list[FolderModel]:
"""
Search for folders for a user where the name matches any of the queries, treating _ and space as equivalent, case-insensitive.
@ -330,16 +336,18 @@ class FolderTable:
return []
results = {}
with get_db_context(db) as db:
folders = db.query(Folder).filter_by(user_id=user_id).all()
async with get_async_db_context(db) as db:
result = await db.execute(select(Folder).filter_by(user_id=user_id))
folders = result.scalars().all()
for folder in folders:
if self.normalize_folder_name(folder.name) in normalized_queries:
results[folder.id] = FolderModel.model_validate(folder)
# get children folders
children = self.get_children_folders_by_id_and_user_id(folder.id, user_id, db=db)
for child in children:
results[child.id] = child
children = await self.get_children_folders_by_id_and_user_id(folder.id, user_id, db=db)
if children:
for child in children:
results[child.id] = child
# Return the results as a list
if not results:
@ -348,16 +356,17 @@ class FolderTable:
results = list(results.values())
return results
def search_folders_by_name_contains(
self, user_id: str, query: str, db: Optional[Session] = None
async def search_folders_by_name_contains(
self, user_id: str, query: str, db: Optional[AsyncSession] = None
) -> list[FolderModel]:
"""
Partial match: normalized name contains (as substring) the normalized query.
"""
normalized_query = self.normalize_folder_name(query)
results = []
with get_db_context(db) as db:
folders = db.query(Folder).filter_by(user_id=user_id).all()
async with get_async_db_context(db) as db:
result = await db.execute(select(Folder).filter_by(user_id=user_id))
folders = result.scalars().all()
for folder in folders:
norm_name = self.normalize_folder_name(folder.name)
if normalized_query in norm_name:

View file

@ -2,8 +2,9 @@ import logging
import time
from typing import Optional
from sqlalchemy.orm import Session, defer
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
from sqlalchemy import select, delete, update
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import Base, JSONField, get_async_db_context
from open_webui.models.users import Users, UserModel, UserResponse
from pydantic import BaseModel, ConfigDict
from sqlalchemy import BigInteger, Boolean, Column, String, Text, Index
@ -107,12 +108,12 @@ class FunctionValves(BaseModel):
class FunctionsTable:
def insert_new_function(
async def insert_new_function(
self,
user_id: str,
type: str,
form_data: FunctionForm,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[FunctionModel]:
function = FunctionModel(
**{
@ -125,11 +126,11 @@ class FunctionsTable:
)
try:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
result = Function(**function.model_dump())
db.add(result)
db.commit()
db.refresh(result)
await db.commit()
await db.refresh(result)
if result:
return FunctionModel.model_validate(result)
else:
@ -138,17 +139,18 @@ class FunctionsTable:
log.exception(f'Error creating a new function: {e}')
return None
def sync_functions(
async def sync_functions(
self,
user_id: str,
functions: list[FunctionWithValvesModel],
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> list[FunctionWithValvesModel]:
# Synchronize functions for a user by updating existing ones, inserting new ones, and removing those that are no longer present.
try:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
# Get existing functions
existing_functions = db.query(Function).all()
result = await db.execute(select(Function))
existing_functions = result.scalars().all()
existing_ids = {func.id for func in existing_functions}
# Prepare a set of new function IDs
@ -157,12 +159,12 @@ class FunctionsTable:
# Update or insert functions
for func in functions:
if func.id in existing_ids:
db.query(Function).filter_by(id=func.id).update(
{
await db.execute(
update(Function).filter_by(id=func.id).values(
**func.model_dump(),
'user_id': user_id,
'updated_at': int(time.time()),
}
user_id=user_id,
updated_at=int(time.time()),
)
)
else:
new_func = Function(
@ -177,24 +179,25 @@ class FunctionsTable:
# Remove functions that are no longer present
for func in existing_functions:
if func.id not in new_function_ids:
db.delete(func)
await db.delete(func)
db.commit()
await db.commit()
return [FunctionModel.model_validate(func) for func in db.query(Function).all()]
result = await db.execute(select(Function))
return [FunctionModel.model_validate(func) for func in result.scalars().all()]
except Exception as e:
log.exception(f'Error syncing functions for user {user_id}: {e}')
return []
def get_function_by_id(self, id: str, db: Optional[Session] = None) -> Optional[FunctionModel]:
async def get_function_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[FunctionModel]:
try:
with get_db_context(db) as db:
function = db.get(Function, id)
return FunctionModel.model_validate(function)
async with get_async_db_context(db) as db:
function = await db.get(Function, id)
return FunctionModel.model_validate(function) if function else None
except Exception:
return None
def get_functions_by_ids(self, ids: list[str], db: Optional[Session] = None) -> list[FunctionModel]:
async def get_functions_by_ids(self, ids: list[str], db: Optional[AsyncSession] = None) -> list[FunctionModel]:
"""
Batch fetch multiple functions by their IDs in a single query.
Returns functions in the same order as the input IDs (None entries filtered out).
@ -202,8 +205,9 @@ class FunctionsTable:
if not ids:
return []
try:
with get_db_context(db) as db:
functions = db.query(Function).filter(Function.id.in_(ids)).all()
async with get_async_db_context(db) as db:
result = await db.execute(select(Function).filter(Function.id.in_(ids)))
functions = result.scalars().all()
# Create a dict for O(1) lookup
func_dict = {f.id: FunctionModel.model_validate(f) for f in functions}
# Return in original order, filtering out any not found
@ -211,27 +215,31 @@ class FunctionsTable:
except Exception:
return []
def get_functions(
self, active_only=False, include_valves=False, db: Optional[Session] = None
async def get_functions(
self, active_only=False, include_valves=False, db: Optional[AsyncSession] = None
) -> list[FunctionModel | FunctionWithValvesModel]:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
if active_only:
functions = db.query(Function).filter_by(is_active=True).all()
result = await db.execute(select(Function).filter_by(is_active=True))
else:
functions = db.query(Function).all()
result = await db.execute(select(Function))
functions = result.scalars().all()
if include_valves:
return [FunctionWithValvesModel.model_validate(function) for function in functions]
else:
return [FunctionModel.model_validate(function) for function in functions]
def get_function_list(self, db: Optional[Session] = None) -> list[FunctionUserResponse]:
with get_db_context(db) as db:
functions = db.query(Function).options(defer(Function.content)).order_by(Function.updated_at.desc()).all()
async def get_function_list(self, db: Optional[AsyncSession] = None) -> list[FunctionUserResponse]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(Function).order_by(Function.updated_at.desc())
)
functions = result.scalars().all()
user_ids = list(set(func.user_id for func in functions))
users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
users = await Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
users_dict = {user.id: user for user in users}
return [
@ -253,42 +261,34 @@ class FunctionsTable:
for func in functions
]
def get_functions_by_type(self, type: str, active_only=False, db: Optional[Session] = None) -> list[FunctionModel]:
with get_db_context(db) as db:
async def get_functions_by_type(self, type: str, active_only=False, db: Optional[AsyncSession] = None) -> list[FunctionModel]:
async with get_async_db_context(db) as db:
if active_only:
return [
FunctionModel.model_validate(function)
for function in db.query(Function).filter_by(type=type, is_active=True).all()
]
result = await db.execute(select(Function).filter_by(type=type, is_active=True))
else:
return [
FunctionModel.model_validate(function) for function in db.query(Function).filter_by(type=type).all()
]
result = await db.execute(select(Function).filter_by(type=type))
return [FunctionModel.model_validate(function) for function in result.scalars().all()]
def get_global_filter_functions(self, db: Optional[Session] = None) -> list[FunctionModel]:
with get_db_context(db) as db:
return [
FunctionModel.model_validate(function)
for function in db.query(Function).filter_by(type='filter', is_active=True, is_global=True).all()
]
async def get_global_filter_functions(self, db: Optional[AsyncSession] = None) -> list[FunctionModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(Function).filter_by(type='filter', is_active=True, is_global=True))
return [FunctionModel.model_validate(function) for function in result.scalars().all()]
def get_global_action_functions(self, db: Optional[Session] = None) -> list[FunctionModel]:
with get_db_context(db) as db:
return [
FunctionModel.model_validate(function)
for function in db.query(Function).filter_by(type='action', is_active=True, is_global=True).all()
]
async def get_global_action_functions(self, db: Optional[AsyncSession] = None) -> list[FunctionModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(Function).filter_by(type='action', is_active=True, is_global=True))
return [FunctionModel.model_validate(function) for function in result.scalars().all()]
def get_function_valves_by_id(self, id: str, db: Optional[Session] = None) -> Optional[dict]:
with get_db_context(db) as db:
async def get_function_valves_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[dict]:
async with get_async_db_context(db) as db:
try:
function = db.get(Function, id)
function = await db.get(Function, id)
return function.valves if function.valves else {}
except Exception as e:
log.exception(f'Error getting function valves by id {id}: {e}')
return None
def get_function_valves_by_ids(self, ids: list[str], db: Optional[Session] = None) -> dict[str, dict]:
async def get_function_valves_by_ids(self, ids: list[str], db: Optional[AsyncSession] = None) -> dict[str, dict]:
"""
Batch fetch valves for multiple functions in a single query.
Returns a dict mapping function_id -> valves dict.
@ -297,33 +297,34 @@ class FunctionsTable:
if not ids:
return {}
try:
with get_db_context(db) as db:
functions = db.query(Function.id, Function.valves).filter(Function.id.in_(ids)).all()
async with get_async_db_context(db) as db:
result = await db.execute(select(Function.id, Function.valves).filter(Function.id.in_(ids)))
functions = result.all()
return {f.id: (f.valves if f.valves else {}) for f in functions}
except Exception as e:
log.exception(f'Error batch-fetching function valves: {e}')
return {}
def update_function_valves_by_id(
self, id: str, valves: dict, db: Optional[Session] = None
async def update_function_valves_by_id(
self, id: str, valves: dict, db: Optional[AsyncSession] = None
) -> Optional[FunctionValves]:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
try:
function = db.get(Function, id)
function = await db.get(Function, id)
function.valves = valves
function.updated_at = int(time.time())
db.commit()
db.refresh(function)
await db.commit()
await db.refresh(function)
return FunctionModel.model_validate(function)
except Exception:
return None
def update_function_metadata_by_id(
self, id: str, metadata: dict, db: Optional[Session] = None
async def update_function_metadata_by_id(
self, id: str, metadata: dict, db: Optional[AsyncSession] = None
) -> Optional[FunctionModel]:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
try:
function = db.get(Function, id)
function = await db.get(Function, id)
if function:
if function.meta:
@ -332,8 +333,8 @@ class FunctionsTable:
function.meta = metadata
function.updated_at = int(time.time())
db.commit()
db.refresh(function)
await db.commit()
await db.refresh(function)
return FunctionModel.model_validate(function)
else:
return None
@ -341,9 +342,9 @@ class FunctionsTable:
log.exception(f'Error updating function metadata by id {id}: {e}')
return None
def get_user_valves_by_id_and_user_id(self, id: str, user_id: str, db: Optional[Session] = None) -> Optional[dict]:
async def get_user_valves_by_id_and_user_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> Optional[dict]:
try:
user = Users.get_user_by_id(user_id, db=db)
user = await Users.get_user_by_id(user_id, db=db)
user_settings = user.settings.model_dump() if user.settings else {}
# Check if user has "functions" and "valves" settings
@ -357,11 +358,11 @@ class FunctionsTable:
log.exception(f'Error getting user values by id {id} and user id {user_id}')
return None
def update_user_valves_by_id_and_user_id(
self, id: str, user_id: str, valves: dict, db: Optional[Session] = None
async def update_user_valves_by_id_and_user_id(
self, id: str, user_id: str, valves: dict, db: Optional[AsyncSession] = None
) -> Optional[dict]:
try:
user = Users.get_user_by_id(user_id, db=db)
user = await Users.get_user_by_id(user_id, db=db)
user_settings = user.settings.model_dump() if user.settings else {}
# Check if user has "functions" and "valves" settings
@ -373,47 +374,47 @@ class FunctionsTable:
user_settings['functions']['valves'][id] = valves
# Update the user settings in the database
Users.update_user_by_id(user_id, {'settings': user_settings}, db=db)
await Users.update_user_by_id(user_id, {'settings': user_settings}, db=db)
return user_settings['functions']['valves'][id]
except Exception as e:
log.exception(f'Error updating user valves by id {id} and user_id {user_id}: {e}')
return None
def update_function_by_id(self, id: str, updated: dict, db: Optional[Session] = None) -> Optional[FunctionModel]:
with get_db_context(db) as db:
async def update_function_by_id(self, id: str, updated: dict, db: Optional[AsyncSession] = None) -> Optional[FunctionModel]:
async with get_async_db_context(db) as db:
try:
db.query(Function).filter_by(id=id).update(
{
await db.execute(
update(Function).filter_by(id=id).values(
**updated,
'updated_at': int(time.time()),
}
updated_at=int(time.time()),
)
)
db.commit()
function = db.get(Function, id)
await db.commit()
function = await db.get(Function, id)
return FunctionModel.model_validate(function) if function else None
except Exception:
return None
def deactivate_all_functions(self, db: Optional[Session] = None) -> Optional[bool]:
with get_db_context(db) as db:
async def deactivate_all_functions(self, db: Optional[AsyncSession] = None) -> Optional[bool]:
async with get_async_db_context(db) as db:
try:
db.query(Function).update(
{
'is_active': False,
'updated_at': int(time.time()),
}
await db.execute(
update(Function).values(
is_active=False,
updated_at=int(time.time()),
)
)
db.commit()
await db.commit()
return True
except Exception:
return None
def delete_function_by_id(self, id: str, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
async def delete_function_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
try:
db.query(Function).filter_by(id=id).delete()
db.commit()
await db.execute(delete(Function).filter_by(id=id))
await db.commit()
return True
except Exception:

View file

@ -4,8 +4,9 @@ import time
from typing import Optional
import uuid
from sqlalchemy.orm import Session
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
from sqlalchemy import select, delete, update, func, and_, or_, cast, String
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import Base, JSONField, get_async_db_context
from open_webui.env import DEFAULT_GROUP_SHARE_PERMISSION
from open_webui.models.files import FileMetadataResponse
@ -15,15 +16,9 @@ from pydantic import BaseModel, ConfigDict
from sqlalchemy import (
BigInteger,
Column,
String,
Text,
JSON,
and_,
func,
ForeignKey,
cast,
or_,
select,
)
log = logging.getLogger(__name__)
@ -143,10 +138,10 @@ class GroupTable:
group_data['data']['config']['share'] = DEFAULT_GROUP_SHARE_PERMISSION
return group_data
def insert_new_group(
self, user_id: str, form_data: GroupForm, db: Optional[Session] = None
async def insert_new_group(
self, user_id: str, form_data: GroupForm, db: Optional[AsyncSession] = None
) -> Optional[GroupModel]:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
group_data = self._ensure_default_share_config(form_data.model_dump(exclude_none=True))
group = GroupModel(
**{
@ -161,8 +156,8 @@ class GroupTable:
try:
result = Group(**group.model_dump())
db.add(result)
db.commit()
db.refresh(result)
await db.commit()
await db.refresh(result)
if result:
return GroupModel.model_validate(result)
else:
@ -171,18 +166,20 @@ class GroupTable:
except Exception:
return None
def get_all_groups(self, db: Optional[Session] = None) -> list[GroupModel]:
with get_db_context(db) as db:
groups = db.query(Group).order_by(Group.updated_at.desc()).all()
async def get_all_groups(self, db: Optional[AsyncSession] = None) -> list[GroupModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(Group).order_by(Group.updated_at.desc()))
groups = result.scalars().all()
return [GroupModel.model_validate(group) for group in groups]
def get_group_by_name(self, name: str, db: Optional[Session] = None) -> Optional[GroupModel]:
with get_db_context(db) as db:
group = db.query(Group).filter(Group.name == name).first()
async def get_group_by_name(self, name: str, db: Optional[AsyncSession] = None) -> Optional[GroupModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(Group).filter(Group.name == name))
group = result.scalars().first()
return GroupModel.model_validate(group) if group else None
def get_groups(self, filter, db: Optional[Session] = None) -> list[GroupResponse]:
with get_db_context(db) as db:
async def get_groups(self, filter, db: Optional[AsyncSession] = None) -> list[GroupResponse]:
async with get_async_db_context(db) as db:
member_count = (
select(func.count(GroupMember.user_id))
.where(GroupMember.group_id == Group.id)
@ -190,11 +187,11 @@ class GroupTable:
.scalar_subquery()
.label('member_count')
)
query = db.query(Group, member_count)
stmt = select(Group, member_count)
if filter:
if 'query' in filter:
query = query.filter(Group.name.ilike(f'%{filter["query"]}%'))
stmt = stmt.filter(Group.name.ilike(f'%{filter["query"]}%'))
# When share filter is present, member check is handled in the share logic
if 'share' in filter:
@ -218,20 +215,21 @@ class GroupTable:
json_share_lower == 'members',
Group.id.in_(member_groups_select),
)
query = query.filter(or_(anyone_can_share, members_only_and_is_member))
stmt = stmt.filter(or_(anyone_can_share, members_only_and_is_member))
else:
query = query.filter(anyone_can_share)
stmt = stmt.filter(anyone_can_share)
else:
query = query.filter(and_(Group.data.isnot(None), json_share_lower == 'false'))
stmt = stmt.filter(and_(Group.data.isnot(None), json_share_lower == 'false'))
else:
# Only apply member_id filter when share filter is NOT present
if 'member_id' in filter:
query = query.filter(
stmt = stmt.filter(
Group.id.in_(select(GroupMember.group_id).where(GroupMember.user_id == filter['member_id']))
)
results = query.order_by(Group.updated_at.desc()).all()
result = await db.execute(stmt.order_by(Group.updated_at.desc()))
rows = result.all()
return [
GroupResponse.model_validate(
@ -240,32 +238,36 @@ class GroupTable:
'member_count': count or 0,
}
)
for group, count in results
for group, count in rows
]
def search_groups(
async def search_groups(
self,
filter: Optional[dict] = None,
skip: int = 0,
limit: int = 30,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> GroupListResponse:
with get_db_context(db) as db:
query = db.query(Group)
async with get_async_db_context(db) as db:
stmt = select(Group)
if filter:
if 'query' in filter:
query = query.filter(Group.name.ilike(f'%{filter["query"]}%'))
stmt = stmt.filter(Group.name.ilike(f'%{filter["query"]}%'))
if 'member_id' in filter:
query = query.filter(
stmt = stmt.filter(
Group.id.in_(select(GroupMember.group_id).where(GroupMember.user_id == filter['member_id']))
)
if 'share' in filter:
share_value = filter['share']
query = query.filter(Group.data.op('->>')('share') == str(share_value))
stmt = stmt.filter(Group.data.op('->>') ('share') == str(share_value))
total = query.count()
# Get total count
count_result = await db.execute(
select(func.count()).select_from(stmt.subquery())
)
total = count_result.scalar()
member_count = (
select(func.count(GroupMember.user_id))
@ -274,7 +276,14 @@ class GroupTable:
.scalar_subquery()
.label('member_count')
)
results = query.add_columns(member_count).order_by(Group.updated_at.desc()).offset(skip).limit(limit).all()
result = await db.execute(
select(Group, member_count)
.where(Group.id.in_(select(stmt.subquery().c.id)))
.order_by(Group.updated_at.desc())
.offset(skip)
.limit(limit)
)
rows = result.all()
return {
'items': [
@ -284,65 +293,67 @@ class GroupTable:
'member_count': count or 0,
}
)
for group, count in results
for group, count in rows
],
'total': total,
}
def get_groups_by_member_id(self, user_id: str, db: Optional[Session] = None) -> list[GroupModel]:
with get_db_context(db) as db:
return [
GroupModel.model_validate(group)
for group in db.query(Group)
async def get_groups_by_member_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[GroupModel]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(Group)
.join(GroupMember, GroupMember.group_id == Group.id)
.filter(GroupMember.user_id == user_id)
.order_by(Group.updated_at.desc())
.all()
]
)
return [GroupModel.model_validate(group) for group in result.scalars().all()]
def get_groups_by_member_ids(
self, user_ids: list[str], db: Optional[Session] = None
async def get_groups_by_member_ids(
self, user_ids: list[str], db: Optional[AsyncSession] = None
) -> dict[str, list[GroupModel]]:
"""Fetch groups for multiple users in a single query to avoid N+1."""
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
# Query GroupMember joined with Group, filtering by user_ids
results = (
db.query(GroupMember.user_id, Group)
result = await db.execute(
select(GroupMember.user_id, Group)
.join(Group, Group.id == GroupMember.group_id)
.filter(GroupMember.user_id.in_(user_ids))
.order_by(Group.updated_at.desc())
.all()
)
rows = result.all()
# Group groups by user_id
user_groups: dict[str, list[GroupModel]] = {uid: [] for uid in user_ids}
for user_id, group in results:
for user_id, group in rows:
user_groups[user_id].append(GroupModel.model_validate(group))
return user_groups
def get_group_by_id(self, id: str, db: Optional[Session] = None) -> Optional[GroupModel]:
async def get_group_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[GroupModel]:
try:
with get_db_context(db) as db:
group = db.query(Group).filter_by(id=id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(Group).filter_by(id=id))
group = result.scalars().first()
return GroupModel.model_validate(group) if group else None
except Exception:
return None
def get_group_user_ids_by_id(self, id: str, db: Optional[Session] = None) -> list[str]:
with get_db_context(db) as db:
members = db.query(GroupMember.user_id).filter(GroupMember.group_id == id).all()
async def get_group_user_ids_by_id(self, id: str, db: Optional[AsyncSession] = None) -> list[str]:
async with get_async_db_context(db) as db:
result = await db.execute(select(GroupMember.user_id).filter(GroupMember.group_id == id))
members = result.all()
if not members:
return []
return [m[0] for m in members]
def get_group_user_ids_by_ids(self, group_ids: list[str], db: Optional[Session] = None) -> dict[str, list[str]]:
with get_db_context(db) as db:
members = (
db.query(GroupMember.group_id, GroupMember.user_id).filter(GroupMember.group_id.in_(group_ids)).all()
async def get_group_user_ids_by_ids(self, group_ids: list[str], db: Optional[AsyncSession] = None) -> dict[str, list[str]]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(GroupMember.group_id, GroupMember.user_id).filter(GroupMember.group_id.in_(group_ids))
)
members = result.all()
group_user_ids: dict[str, list[str]] = {group_id: [] for group_id in group_ids}
@ -351,10 +362,10 @@ class GroupTable:
return group_user_ids
def set_group_user_ids_by_id(self, group_id: str, user_ids: list[str], db: Optional[Session] = None) -> None:
with get_db_context(db) as db:
async def set_group_user_ids_by_id(self, group_id: str, user_ids: list[str], db: Optional[AsyncSession] = None) -> None:
async with get_async_db_context(db) as db:
# Delete existing members
db.query(GroupMember).filter(GroupMember.group_id == group_id).delete()
await db.execute(delete(GroupMember).filter(GroupMember.group_id == group_id))
# Insert new members
now = int(time.time())
@ -370,101 +381,106 @@ class GroupTable:
]
db.add_all(new_members)
db.commit()
await db.commit()
def get_group_member_count_by_id(self, id: str, db: Optional[Session] = None) -> int:
with get_db_context(db) as db:
count = db.query(func.count(GroupMember.user_id)).filter(GroupMember.group_id == id).scalar()
async def get_group_member_count_by_id(self, id: str, db: Optional[AsyncSession] = None) -> int:
async with get_async_db_context(db) as db:
result = await db.execute(select(func.count(GroupMember.user_id)).filter(GroupMember.group_id == id))
count = result.scalar()
return count if count else 0
def get_group_member_counts_by_ids(self, ids: list[str], db: Optional[Session] = None) -> dict[str, int]:
async def get_group_member_counts_by_ids(self, ids: list[str], db: Optional[AsyncSession] = None) -> dict[str, int]:
if not ids:
return {}
with get_db_context(db) as db:
rows = (
db.query(GroupMember.group_id, func.count(GroupMember.user_id))
async with get_async_db_context(db) as db:
result = await db.execute(
select(GroupMember.group_id, func.count(GroupMember.user_id))
.filter(GroupMember.group_id.in_(ids))
.group_by(GroupMember.group_id)
.all()
)
rows = result.all()
return {group_id: count for group_id, count in rows}
def update_group_by_id(
async def update_group_by_id(
self,
id: str,
form_data: GroupUpdateForm,
overwrite: bool = False,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[GroupModel]:
try:
with get_db_context(db) as db:
db.query(Group).filter_by(id=id).update(
{
async with get_async_db_context(db) as db:
await db.execute(
update(Group).filter_by(id=id).values(
**form_data.model_dump(exclude_none=True),
'updated_at': int(time.time()),
}
updated_at=int(time.time()),
)
)
db.commit()
return self.get_group_by_id(id=id, db=db)
await db.commit()
return await self.get_group_by_id(id=id, db=db)
except Exception as e:
log.exception(e)
return None
def delete_group_by_id(self, id: str, db: Optional[Session] = None) -> bool:
async def delete_group_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
try:
with get_db_context(db) as db:
db.query(Group).filter_by(id=id).delete()
db.commit()
async with get_async_db_context(db) as db:
await db.execute(delete(Group).filter_by(id=id))
await db.commit()
return True
except Exception:
return False
def delete_all_groups(self, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
async def delete_all_groups(self, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
try:
db.query(Group).delete()
db.commit()
await db.execute(delete(Group))
await db.commit()
return True
except Exception:
return False
def remove_user_from_all_groups(self, user_id: str, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
async def remove_user_from_all_groups(self, user_id: str, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
try:
# Find all groups the user belongs to
groups = (
db.query(Group)
result = await db.execute(
select(Group)
.join(GroupMember, GroupMember.group_id == Group.id)
.filter(GroupMember.user_id == user_id)
.all()
)
groups = result.scalars().all()
# Remove the user from each group
for group in groups:
db.query(GroupMember).filter(
GroupMember.group_id == group.id, GroupMember.user_id == user_id
).delete()
await db.execute(
delete(GroupMember).filter(
GroupMember.group_id == group.id, GroupMember.user_id == user_id
)
)
db.query(Group).filter_by(id=group.id).update({'updated_at': int(time.time())})
await db.execute(
update(Group).filter_by(id=group.id).values(updated_at=int(time.time()))
)
db.commit()
await db.commit()
return True
except Exception:
db.rollback()
await db.rollback()
return False
def create_groups_by_group_names(
self, user_id: str, group_names: list[str], db: Optional[Session] = None
async def create_groups_by_group_names(
self, user_id: str, group_names: list[str], db: Optional[AsyncSession] = None
) -> list[GroupModel]:
# check for existing groups
existing_groups = self.get_all_groups(db=db)
existing_groups = await self.get_all_groups(db=db)
existing_group_names = {group.name for group in existing_groups}
new_groups = []
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
for group_name in group_names:
if group_name not in existing_group_names:
new_group = GroupModel(
@ -483,31 +499,31 @@ class GroupTable:
try:
result = Group(**new_group.model_dump())
db.add(result)
db.commit()
db.refresh(result)
await db.commit()
await db.refresh(result)
new_groups.append(GroupModel.model_validate(result))
except Exception as e:
log.exception(e)
continue
return new_groups
def sync_groups_by_group_names(self, user_id: str, group_names: list[str], db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
async def sync_groups_by_group_names(self, user_id: str, group_names: list[str], db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
try:
now = int(time.time())
# 1. Groups that SHOULD contain the user
target_groups = db.query(Group).filter(Group.name.in_(group_names)).all()
result = await db.execute(select(Group).filter(Group.name.in_(group_names)))
target_groups = result.scalars().all()
target_group_ids = {g.id for g in target_groups}
# 2. Groups the user is CURRENTLY in
existing_group_ids = {
g.id
for g in db.query(Group)
result = await db.execute(
select(Group)
.join(GroupMember, GroupMember.group_id == Group.id)
.filter(GroupMember.user_id == user_id)
.all()
}
)
existing_group_ids = {g.id for g in result.scalars().all()}
# 3. Determine adds + removals
groups_to_add = target_group_ids - existing_group_ids
@ -515,13 +531,15 @@ class GroupTable:
# 4. Remove in one bulk delete
if groups_to_remove:
db.query(GroupMember).filter(
GroupMember.user_id == user_id,
GroupMember.group_id.in_(groups_to_remove),
).delete(synchronize_session=False)
await db.execute(
delete(GroupMember).filter(
GroupMember.user_id == user_id,
GroupMember.group_id.in_(groups_to_remove),
)
)
db.query(Group).filter(Group.id.in_(groups_to_remove)).update(
{'updated_at': now}, synchronize_session=False
await db.execute(
update(Group).filter(Group.id.in_(groups_to_remove)).values(updated_at=now)
)
# 5. Bulk insert missing memberships
@ -537,27 +555,28 @@ class GroupTable:
)
if groups_to_add:
db.query(Group).filter(Group.id.in_(groups_to_add)).update(
{'updated_at': now}, synchronize_session=False
await db.execute(
update(Group).filter(Group.id.in_(groups_to_add)).values(updated_at=now)
)
db.commit()
await db.commit()
return True
except Exception as e:
log.exception(e)
db.rollback()
await db.rollback()
return False
def add_users_to_group(
async def add_users_to_group(
self,
id: str,
user_ids: Optional[list[str]] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[GroupModel]:
try:
with get_db_context(db) as db:
group = db.query(Group).filter_by(id=id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(Group).filter_by(id=id))
group = result.scalars().first()
if not group:
return None
@ -574,15 +593,14 @@ class GroupTable:
updated_at=now,
)
)
db.flush() # Detect unique constraint violation early
await db.flush() # Detect unique constraint violation early
except Exception:
db.rollback() # Clear failed INSERT
db.begin() # Start a new transaction
await db.rollback() # Clear failed INSERT
continue # Duplicate → ignore
group.updated_at = now
db.commit()
db.refresh(group)
await db.commit()
await db.refresh(group)
return GroupModel.model_validate(group)
@ -590,15 +608,16 @@ class GroupTable:
log.exception(e)
return None
def remove_users_from_group(
async def remove_users_from_group(
self,
id: str,
user_ids: Optional[list[str]] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[GroupModel]:
try:
with get_db_context(db) as db:
group = db.query(Group).filter_by(id=id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(Group).filter_by(id=id))
group = result.scalars().first()
if not group:
return None
@ -606,15 +625,15 @@ class GroupTable:
return GroupModel.model_validate(group)
# Remove users from group_member in batch
db.query(GroupMember).filter(GroupMember.group_id == id, GroupMember.user_id.in_(user_ids)).delete(
synchronize_session=False
await db.execute(
delete(GroupMember).filter(GroupMember.group_id == id, GroupMember.user_id.in_(user_ids))
)
# Update group timestamp
group.updated_at = int(time.time())
db.commit()
db.refresh(group)
await db.commit()
await db.refresh(group)
return GroupModel.model_validate(group)
except Exception as e:

View file

@ -4,8 +4,9 @@ import time
from typing import Optional
import uuid
from sqlalchemy.orm import Session
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
from sqlalchemy import select, delete, update, or_, func
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import Base, JSONField, get_async_db_context
from open_webui.models.files import (
File,
@ -27,7 +28,6 @@ from sqlalchemy import (
Text,
JSON,
UniqueConstraint,
or_,
)
log = logging.getLogger(__name__)
@ -134,25 +134,25 @@ class KnowledgeFileListResponse(BaseModel):
class KnowledgeTable:
def _get_access_grants(self, knowledge_id: str, db: Optional[Session] = None) -> list[AccessGrantModel]:
return AccessGrants.get_grants_by_resource('knowledge', knowledge_id, db=db)
async def _get_access_grants(self, knowledge_id: str, db: Optional[AsyncSession] = None) -> list[AccessGrantModel]:
return await AccessGrants.get_grants_by_resource('knowledge', knowledge_id, db=db)
def _to_knowledge_model(
async def _to_knowledge_model(
self,
knowledge: Knowledge,
access_grants: Optional[list[AccessGrantModel]] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> KnowledgeModel:
knowledge_data = KnowledgeModel.model_validate(knowledge).model_dump(exclude={'access_grants'})
knowledge_data['access_grants'] = (
access_grants if access_grants is not None else self._get_access_grants(knowledge_data['id'], db=db)
access_grants if access_grants is not None else await self._get_access_grants(knowledge_data['id'], db=db)
)
return KnowledgeModel.model_validate(knowledge_data)
def insert_new_knowledge(
self, user_id: str, form_data: KnowledgeForm, db: Optional[Session] = None
async def insert_new_knowledge(
self, user_id: str, form_data: KnowledgeForm, db: Optional[AsyncSession] = None
) -> Optional[KnowledgeModel]:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
knowledge = KnowledgeModel(
**{
**form_data.model_dump(exclude={'access_grants'}),
@ -167,27 +167,28 @@ class KnowledgeTable:
try:
result = Knowledge(**knowledge.model_dump(exclude={'access_grants'}))
db.add(result)
db.commit()
db.refresh(result)
AccessGrants.set_access_grants('knowledge', result.id, form_data.access_grants, db=db)
await db.commit()
await db.refresh(result)
await AccessGrants.set_access_grants('knowledge', result.id, form_data.access_grants, db=db)
if result:
return self._to_knowledge_model(result, db=db)
return await self._to_knowledge_model(result, db=db)
else:
return None
except Exception:
return None
def get_knowledge_bases(
self, skip: int = 0, limit: int = 30, db: Optional[Session] = None
async def get_knowledge_bases(
self, skip: int = 0, limit: int = 30, db: Optional[AsyncSession] = None
) -> list[KnowledgeUserModel]:
with get_db_context(db) as db:
all_knowledge = db.query(Knowledge).order_by(Knowledge.updated_at.desc()).all()
async with get_async_db_context(db) as db:
result = await db.execute(select(Knowledge).order_by(Knowledge.updated_at.desc()))
all_knowledge = result.scalars().all()
user_ids = list(set(knowledge.user_id for knowledge in all_knowledge))
knowledge_ids = [knowledge.id for knowledge in all_knowledge]
users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
users = await Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
users_dict = {user.id: user for user in users}
grants_map = AccessGrants.get_grants_by_resources('knowledge', knowledge_ids, db=db)
grants_map = await AccessGrants.get_grants_by_resources('knowledge', knowledge_ids, db=db)
knowledge_bases = []
for knowledge in all_knowledge:
@ -195,33 +196,33 @@ class KnowledgeTable:
knowledge_bases.append(
KnowledgeUserModel.model_validate(
{
**self._to_knowledge_model(
**(await self._to_knowledge_model(
knowledge,
access_grants=grants_map.get(knowledge.id, []),
db=db,
).model_dump(),
)).model_dump(),
'user': user.model_dump() if user else None,
}
)
)
return knowledge_bases
def search_knowledge_bases(
async def search_knowledge_bases(
self,
user_id: str,
filter: dict,
skip: int = 0,
limit: int = 30,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> KnowledgeListResponse:
try:
with get_db_context(db) as db:
query = db.query(Knowledge, User).outerjoin(User, User.id == Knowledge.user_id)
async with get_async_db_context(db) as db:
stmt = select(Knowledge, User).outerjoin(User, User.id == Knowledge.user_id)
if filter:
query_key = filter.get('query')
if query_key:
query = query.filter(
stmt = stmt.filter(
or_(
Knowledge.name.ilike(f'%{query_key}%'),
Knowledge.description.ilike(f'%{query_key}%'),
@ -233,42 +234,46 @@ class KnowledgeTable:
view_option = filter.get('view_option')
if view_option == 'created':
query = query.filter(Knowledge.user_id == user_id)
stmt = stmt.filter(Knowledge.user_id == user_id)
elif view_option == 'shared':
query = query.filter(Knowledge.user_id != user_id)
stmt = stmt.filter(Knowledge.user_id != user_id)
query = AccessGrants.has_permission_filter(
stmt = AccessGrants.has_permission_filter(
db=db,
query=query,
query=stmt,
DocumentModel=Knowledge,
filter=filter,
resource_type='knowledge',
permission='read',
)
query = query.order_by(Knowledge.updated_at.desc(), Knowledge.id.asc())
stmt = stmt.order_by(Knowledge.updated_at.desc(), Knowledge.id.asc())
total = query.count()
count_result = await db.execute(
select(func.count()).select_from(stmt.subquery())
)
total = count_result.scalar()
if skip:
query = query.offset(skip)
stmt = stmt.offset(skip)
if limit:
query = query.limit(limit)
stmt = stmt.limit(limit)
items = query.all()
result = await db.execute(stmt)
items = result.all()
knowledge_ids = [kb.id for kb, _ in items]
grants_map = AccessGrants.get_grants_by_resources('knowledge', knowledge_ids, db=db)
grants_map = await AccessGrants.get_grants_by_resources('knowledge', knowledge_ids, db=db)
knowledge_bases = []
for knowledge_base, user in items:
knowledge_bases.append(
KnowledgeUserModel.model_validate(
{
**self._to_knowledge_model(
**(await self._to_knowledge_model(
knowledge_base,
access_grants=grants_map.get(knowledge_base.id, []),
db=db,
).model_dump(),
)).model_dump(),
'user': (UserModel.model_validate(user).model_dump() if user else None),
}
)
@ -279,28 +284,27 @@ class KnowledgeTable:
print(e)
return KnowledgeListResponse(items=[], total=0)
def search_knowledge_files(
self, filter: dict, skip: int = 0, limit: int = 30, db: Optional[Session] = None
async def search_knowledge_files(
self, filter: dict, skip: int = 0, limit: int = 30, db: Optional[AsyncSession] = None
) -> KnowledgeFileListResponse:
"""
Scalable version: search files across all knowledge bases the user has
READ access to, without loading all KBs or using large IN() lists.
"""
try:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
# Base query: join Knowledge → KnowledgeFile → File
query = (
db.query(File, User, Knowledge)
stmt = (
select(File, User, Knowledge)
.join(KnowledgeFile, File.id == KnowledgeFile.file_id)
.join(Knowledge, KnowledgeFile.knowledge_id == Knowledge.id)
.outerjoin(User, User.id == KnowledgeFile.user_id)
)
# Apply access-control directly to the joined query
# This makes the database handle filtering, even with 10k+ KBs
query = AccessGrants.has_permission_filter(
stmt = AccessGrants.has_permission_filter(
db=db,
query=query,
query=stmt,
DocumentModel=Knowledge,
filter=filter,
resource_type='knowledge',
@ -311,20 +315,24 @@ class KnowledgeTable:
if filter:
q = filter.get('query')
if q:
query = query.filter(File.filename.ilike(f'%{q}%'))
stmt = stmt.filter(File.filename.ilike(f'%{q}%'))
# Order by file changes
query = query.order_by(File.updated_at.desc(), File.id.asc())
stmt = stmt.order_by(File.updated_at.desc(), File.id.asc())
# Count before pagination
total = query.count()
count_result = await db.execute(
select(func.count()).select_from(stmt.subquery())
)
total = count_result.scalar()
if skip:
query = query.offset(skip)
stmt = stmt.offset(skip)
if limit:
query = query.limit(limit)
stmt = stmt.limit(limit)
rows = query.all()
result = await db.execute(stmt)
rows = result.all()
items = []
for file, user, knowledge in rows:
@ -332,7 +340,7 @@ class KnowledgeTable:
FileUserResponse(
**FileModel.model_validate(file).model_dump(),
user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
collection=self._to_knowledge_model(knowledge, db=db).model_dump(),
collection=(await self._to_knowledge_model(knowledge, db=db)).model_dump(),
)
)
@ -342,14 +350,15 @@ class KnowledgeTable:
print('search_knowledge_files error:', e)
return KnowledgeFileListResponse(items=[], total=0)
def check_access_by_user_id(self, id, user_id, permission='write', db: Optional[Session] = None) -> bool:
knowledge = self.get_knowledge_by_id(id, db=db)
async def check_access_by_user_id(self, id, user_id, permission='write', db: Optional[AsyncSession] = None) -> bool:
knowledge = await self.get_knowledge_by_id(id, db=db)
if not knowledge:
return False
if knowledge.user_id == user_id:
return True
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id, db=db)}
return AccessGrants.has_access(
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
user_group_ids = {group.id for group in user_groups}
return await AccessGrants.has_access(
user_id=user_id,
resource_type='knowledge',
resource_id=knowledge.id,
@ -358,45 +367,50 @@ class KnowledgeTable:
db=db,
)
def get_knowledge_bases_by_user_id(
self, user_id: str, permission: str = 'write', db: Optional[Session] = None
async def get_knowledge_bases_by_user_id(
self, user_id: str, permission: str = 'write', db: Optional[AsyncSession] = None
) -> list[KnowledgeUserModel]:
knowledge_bases = self.get_knowledge_bases(db=db)
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id, db=db)}
return [
knowledge_base
for knowledge_base in knowledge_bases
if knowledge_base.user_id == user_id
or AccessGrants.has_access(
knowledge_bases = await self.get_knowledge_bases(db=db)
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
user_group_ids = {group.id for group in user_groups}
result = []
for knowledge_base in knowledge_bases:
if knowledge_base.user_id == user_id:
result.append(knowledge_base)
elif await AccessGrants.has_access(
user_id=user_id,
resource_type='knowledge',
resource_id=knowledge_base.id,
permission=permission,
user_group_ids=user_group_ids,
db=db,
)
]
):
result.append(knowledge_base)
return result
def get_knowledge_by_id(self, id: str, db: Optional[Session] = None) -> Optional[KnowledgeModel]:
async def get_knowledge_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[KnowledgeModel]:
try:
with get_db_context(db) as db:
knowledge = db.query(Knowledge).filter_by(id=id).first()
return self._to_knowledge_model(knowledge, db=db) if knowledge else None
async with get_async_db_context(db) as db:
result = await db.execute(select(Knowledge).filter_by(id=id))
knowledge = result.scalars().first()
return await self._to_knowledge_model(knowledge, db=db) if knowledge else None
except Exception:
return None
def get_knowledge_by_id_and_user_id(
self, id: str, user_id: str, db: Optional[Session] = None
async def get_knowledge_by_id_and_user_id(
self, id: str, user_id: str, db: Optional[AsyncSession] = None
) -> Optional[KnowledgeModel]:
knowledge = self.get_knowledge_by_id(id, db=db)
knowledge = await self.get_knowledge_by_id(id, db=db)
if not knowledge:
return None
if knowledge.user_id == user_id:
return knowledge
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id, db=db)}
if AccessGrants.has_access(
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
user_group_ids = {group.id for group in user_groups}
if await AccessGrants.has_access(
user_id=user_id,
resource_type='knowledge',
resource_id=knowledge.id,
@ -407,19 +421,19 @@ class KnowledgeTable:
return knowledge
return None
def get_knowledges_by_file_id(self, file_id: str, db: Optional[Session] = None) -> list[KnowledgeModel]:
async def get_knowledges_by_file_id(self, file_id: str, db: Optional[AsyncSession] = None) -> list[KnowledgeModel]:
try:
with get_db_context(db) as db:
knowledges = (
db.query(Knowledge)
async with get_async_db_context(db) as db:
result = await db.execute(
select(Knowledge)
.join(KnowledgeFile, Knowledge.id == KnowledgeFile.knowledge_id)
.filter(KnowledgeFile.file_id == file_id)
.all()
)
knowledges = result.scalars().all()
knowledge_ids = [k.id for k in knowledges]
grants_map = AccessGrants.get_grants_by_resources('knowledge', knowledge_ids, db=db)
grants_map = await AccessGrants.get_grants_by_resources('knowledge', knowledge_ids, db=db)
return [
self._to_knowledge_model(
await self._to_knowledge_model(
knowledge,
access_grants=grants_map.get(knowledge.id, []),
db=db,
@ -429,19 +443,19 @@ class KnowledgeTable:
except Exception:
return []
def search_files_by_id(
async def search_files_by_id(
self,
knowledge_id: str,
user_id: str,
filter: dict,
skip: int = 0,
limit: int = 30,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> KnowledgeFileListResponse:
try:
with get_db_context(db) as db:
query = (
db.query(File, User)
async with get_async_db_context(db) as db:
stmt = (
select(File, User)
.join(KnowledgeFile, File.id == KnowledgeFile.file_id)
.outerjoin(User, User.id == KnowledgeFile.user_id)
.filter(KnowledgeFile.knowledge_id == knowledge_id)
@ -453,13 +467,13 @@ class KnowledgeTable:
if filter:
query_key = filter.get('query')
if query_key:
query = query.filter(or_(File.filename.ilike(f'%{query_key}%')))
stmt = stmt.filter(or_(File.filename.ilike(f'%{query_key}%')))
view_option = filter.get('view_option')
if view_option == 'created':
query = query.filter(KnowledgeFile.user_id == user_id)
stmt = stmt.filter(KnowledgeFile.user_id == user_id)
elif view_option == 'shared':
query = query.filter(KnowledgeFile.user_id != user_id)
stmt = stmt.filter(KnowledgeFile.user_id != user_id)
order_by = filter.get('order_by')
direction = filter.get('direction')
@ -473,17 +487,21 @@ class KnowledgeTable:
primary_sort = File.updated_at.asc() if is_asc else File.updated_at.desc()
# Apply sort with secondary key for deterministic pagination
query = query.order_by(primary_sort, File.id.asc())
stmt = stmt.order_by(primary_sort, File.id.asc())
# Count BEFORE pagination
total = query.count()
count_result = await db.execute(
select(func.count()).select_from(stmt.subquery())
)
total = count_result.scalar()
if skip:
query = query.offset(skip)
stmt = stmt.offset(skip)
if limit:
query = query.limit(limit)
stmt = stmt.limit(limit)
items = query.all()
result = await db.execute(stmt)
items = result.all()
files = []
for file, user in items:
@ -499,35 +517,34 @@ class KnowledgeTable:
print(e)
return KnowledgeFileListResponse(items=[], total=0)
def get_files_by_id(self, knowledge_id: str, db: Optional[Session] = None) -> list[FileModel]:
async def get_files_by_id(self, knowledge_id: str, db: Optional[AsyncSession] = None) -> list[FileModel]:
try:
with get_db_context(db) as db:
files = (
db.query(File)
async with get_async_db_context(db) as db:
result = await db.execute(
select(File)
.join(KnowledgeFile, File.id == KnowledgeFile.file_id)
.filter(KnowledgeFile.knowledge_id == knowledge_id)
.all()
)
files = result.scalars().all()
return [FileModel.model_validate(file) for file in files]
except Exception:
return []
def get_file_metadatas_by_id(self, knowledge_id: str, db: Optional[Session] = None) -> list[FileMetadataResponse]:
async def get_file_metadatas_by_id(self, knowledge_id: str, db: Optional[AsyncSession] = None) -> list[FileMetadataResponse]:
try:
with get_db_context(db) as db:
files = self.get_files_by_id(knowledge_id, db=db)
return [FileMetadataResponse(**file.model_dump()) for file in files]
files = await self.get_files_by_id(knowledge_id, db=db)
return [FileMetadataResponse(**file.model_dump()) for file in files]
except Exception:
return []
def add_file_to_knowledge_by_id(
async def add_file_to_knowledge_by_id(
self,
knowledge_id: str,
file_id: str,
user_id: str,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[KnowledgeFileModel]:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
knowledge_file = KnowledgeFileModel(
**{
'id': str(uuid.uuid4()),
@ -542,8 +559,8 @@ class KnowledgeTable:
try:
result = KnowledgeFile(**knowledge_file.model_dump())
db.add(result)
db.commit()
db.refresh(result)
await db.commit()
await db.refresh(result)
if result:
return KnowledgeFileModel.model_validate(result)
else:
@ -551,103 +568,103 @@ class KnowledgeTable:
except Exception:
return None
def has_file(self, knowledge_id: str, file_id: str, db: Optional[Session] = None) -> bool:
async def has_file(self, knowledge_id: str, file_id: str, db: Optional[AsyncSession] = None) -> bool:
"""Check whether a file belongs to a knowledge base."""
try:
with get_db_context(db) as db:
return db.query(KnowledgeFile).filter_by(knowledge_id=knowledge_id, file_id=file_id).first() is not None
async with get_async_db_context(db) as db:
result = await db.execute(
select(KnowledgeFile).filter_by(knowledge_id=knowledge_id, file_id=file_id).limit(1)
)
return result.scalars().first() is not None
except Exception:
return False
def remove_file_from_knowledge_by_id(self, knowledge_id: str, file_id: str, db: Optional[Session] = None) -> bool:
async def remove_file_from_knowledge_by_id(self, knowledge_id: str, file_id: str, db: Optional[AsyncSession] = None) -> bool:
try:
with get_db_context(db) as db:
db.query(KnowledgeFile).filter_by(knowledge_id=knowledge_id, file_id=file_id).delete()
db.commit()
async with get_async_db_context(db) as db:
await db.execute(delete(KnowledgeFile).filter_by(knowledge_id=knowledge_id, file_id=file_id))
await db.commit()
return True
except Exception:
return False
def reset_knowledge_by_id(self, id: str, db: Optional[Session] = None) -> Optional[KnowledgeModel]:
async def reset_knowledge_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[KnowledgeModel]:
try:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
# Delete all knowledge_file entries for this knowledge_id
db.query(KnowledgeFile).filter_by(knowledge_id=id).delete()
db.commit()
await db.execute(delete(KnowledgeFile).filter_by(knowledge_id=id))
await db.commit()
# Update the knowledge entry's updated_at timestamp
db.query(Knowledge).filter_by(id=id).update(
{
'updated_at': int(time.time()),
}
await db.execute(
update(Knowledge).filter_by(id=id).values(updated_at=int(time.time()))
)
db.commit()
await db.commit()
return self.get_knowledge_by_id(id=id, db=db)
return await self.get_knowledge_by_id(id=id, db=db)
except Exception as e:
log.exception(e)
return None
def update_knowledge_by_id(
async def update_knowledge_by_id(
self,
id: str,
form_data: KnowledgeForm,
overwrite: bool = False,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[KnowledgeModel]:
try:
with get_db_context(db) as db:
knowledge = self.get_knowledge_by_id(id=id, db=db)
db.query(Knowledge).filter_by(id=id).update(
{
async with get_async_db_context(db) as db:
await db.execute(
update(Knowledge).filter_by(id=id).values(
**form_data.model_dump(exclude={'access_grants'}),
'updated_at': int(time.time()),
}
updated_at=int(time.time()),
)
)
db.commit()
await db.commit()
if form_data.access_grants is not None:
AccessGrants.set_access_grants('knowledge', id, form_data.access_grants, db=db)
return self.get_knowledge_by_id(id=id, db=db)
await AccessGrants.set_access_grants('knowledge', id, form_data.access_grants, db=db)
return await self.get_knowledge_by_id(id=id, db=db)
except Exception as e:
log.exception(e)
return None
def update_knowledge_data_by_id(
self, id: str, data: dict, db: Optional[Session] = None
async def update_knowledge_data_by_id(
self, id: str, data: dict, db: Optional[AsyncSession] = None
) -> Optional[KnowledgeModel]:
try:
with get_db_context(db) as db:
knowledge = self.get_knowledge_by_id(id=id, db=db)
db.query(Knowledge).filter_by(id=id).update(
{
'data': data,
'updated_at': int(time.time()),
}
async with get_async_db_context(db) as db:
await db.execute(
update(Knowledge).filter_by(id=id).values(
data=data,
updated_at=int(time.time()),
)
)
db.commit()
return self.get_knowledge_by_id(id=id, db=db)
await db.commit()
return await self.get_knowledge_by_id(id=id, db=db)
except Exception as e:
log.exception(e)
return None
def delete_knowledge_by_id(self, id: str, db: Optional[Session] = None) -> bool:
async def delete_knowledge_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
try:
with get_db_context(db) as db:
AccessGrants.revoke_all_access('knowledge', id, db=db)
db.query(Knowledge).filter_by(id=id).delete()
db.commit()
async with get_async_db_context(db) as db:
await AccessGrants.revoke_all_access('knowledge', id, db=db)
await db.execute(delete(Knowledge).filter_by(id=id))
await db.commit()
return True
except Exception:
return False
def delete_all_knowledge(self, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
async def delete_all_knowledge(self, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
try:
knowledge_ids = [row[0] for row in db.query(Knowledge.id).all()]
result = await db.execute(select(Knowledge.id))
knowledge_ids = [row[0] for row in result.all()]
for knowledge_id in knowledge_ids:
AccessGrants.revoke_all_access('knowledge', knowledge_id, db=db)
db.query(Knowledge).delete()
db.commit()
await AccessGrants.revoke_all_access('knowledge', knowledge_id, db=db)
await db.execute(delete(Knowledge))
await db.commit()
return True
except Exception:

View file

@ -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:

View file

@ -3,8 +3,9 @@ import time
import uuid
from typing import Optional
from sqlalchemy.orm import Session
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
from sqlalchemy import select, delete, func
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import Base, JSONField, get_async_db_context
from open_webui.models.tags import TagModel, Tag, Tags
from open_webui.models.users import Users, User, UserNameResponse
from open_webui.models.channels import Channels, ChannelMember
@ -12,7 +13,7 @@ from open_webui.models.channels import Channels, ChannelMember
from pydantic import BaseModel, ConfigDict, field_validator
from sqlalchemy import BigInteger, Boolean, Column, String, Text, JSON
from sqlalchemy import or_, func, select, and_, text
from sqlalchemy import or_, func, and_, text
from sqlalchemy.sql import exists
####################
@ -137,15 +138,15 @@ class MessageResponse(MessageReplyToResponse):
class MessageTable:
def insert_new_message(
async def insert_new_message(
self,
form_data: MessageForm,
channel_id: str,
user_id: str,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[MessageModel]:
with get_db_context(db) as db:
channel_member = Channels.join_channel(channel_id, user_id)
async with get_async_db_context(db) as db:
channel_member = await Channels.join_channel(channel_id, user_id)
id = str(uuid.uuid4())
ts = int(time.time_ns())
@ -170,38 +171,38 @@ class MessageTable:
result = Message(**message.model_dump())
db.add(result)
db.commit()
db.refresh(result)
await db.commit()
await db.refresh(result)
return MessageModel.model_validate(result) if result else None
def get_message_by_id(
async def get_message_by_id(
self,
id: str,
include_thread_replies: Optional[bool] = True,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[MessageResponse]:
with get_db_context(db) as db:
message = db.get(Message, id)
async with get_async_db_context(db) as db:
message = await db.get(Message, id)
if not message:
return None
reply_to_message = (
self.get_message_by_id(message.reply_to_id, include_thread_replies=False, db=db)
await self.get_message_by_id(message.reply_to_id, include_thread_replies=False, db=db)
if message.reply_to_id
else None
)
reactions = self.get_reactions_by_message_id(id, db=db)
reactions = await self.get_reactions_by_message_id(id, db=db)
thread_replies = []
if include_thread_replies:
thread_replies = self.get_thread_replies_by_message_id(id, db=db)
thread_replies = await self.get_thread_replies_by_message_id(id, db=db)
# Check if message was sent by webhook (webhook info in meta takes precedence)
webhook_info = message.meta.get('webhook') if message.meta else None
if webhook_info and webhook_info.get('id'):
# Look up webhook by ID to get current name
webhook = Channels.get_webhook_by_id(webhook_info.get('id'), db=db)
webhook = await Channels.get_webhook_by_id(webhook_info.get('id'), db=db)
if webhook:
user_info = {
'id': webhook.id,
@ -216,7 +217,7 @@ class MessageTable:
'role': 'webhook',
}
else:
user = Users.get_user_by_id(message.user_id, db=db)
user = await Users.get_user_by_id(message.user_id, db=db)
user_info = user.model_dump() if user else None
return MessageResponse.model_validate(
@ -230,34 +231,41 @@ class MessageTable:
}
)
def get_thread_replies_by_message_id(self, id: str, db: Optional[Session] = None) -> list[MessageReplyToResponse]:
with get_db_context(db) as db:
all_messages = db.query(Message).filter_by(parent_id=id).order_by(Message.created_at.desc()).all()
async def _resolve_user_info(self, message: Message, db: AsyncSession) -> Optional[dict]:
"""Resolve user info from message, handling webhook messages."""
webhook_info = message.meta.get('webhook') if message.meta else None
if webhook_info and webhook_info.get('id'):
webhook = await Channels.get_webhook_by_id(webhook_info.get('id'), db=db)
if webhook:
return {
'id': webhook.id,
'name': webhook.name,
'role': 'webhook',
}
else:
return {
'id': webhook_info.get('id'),
'name': 'Deleted Webhook',
'role': 'webhook',
}
return None
async def get_thread_replies_by_message_id(self, id: str, db: Optional[AsyncSession] = None) -> list[MessageReplyToResponse]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(Message).filter_by(parent_id=id).order_by(Message.created_at.desc())
)
all_messages = result.scalars().all()
messages = []
for message in all_messages:
reply_to_message = (
self.get_message_by_id(message.reply_to_id, include_thread_replies=False, db=db)
await self.get_message_by_id(message.reply_to_id, include_thread_replies=False, db=db)
if message.reply_to_id
else None
)
webhook_info = message.meta.get('webhook') if message.meta else None
user_info = None
if webhook_info and webhook_info.get('id'):
webhook = Channels.get_webhook_by_id(webhook_info.get('id'), db=db)
if webhook:
user_info = {
'id': webhook.id,
'name': webhook.name,
'role': 'webhook',
}
else:
user_info = {
'id': webhook_info.get('id'),
'name': 'Deleted Webhook',
'role': 'webhook',
}
user_info = await self._resolve_user_info(message, db)
messages.append(
MessageReplyToResponse.model_validate(
@ -270,51 +278,37 @@ class MessageTable:
)
return messages
def get_reply_user_ids_by_message_id(self, id: str, db: Optional[Session] = None) -> list[str]:
with get_db_context(db) as db:
return [message.user_id for message in db.query(Message).filter_by(parent_id=id).all()]
async def get_reply_user_ids_by_message_id(self, id: str, db: Optional[AsyncSession] = None) -> list[str]:
async with get_async_db_context(db) as db:
result = await db.execute(select(Message.user_id).filter_by(parent_id=id))
return [row[0] for row in result.all()]
def get_messages_by_channel_id(
async def get_messages_by_channel_id(
self,
channel_id: str,
skip: int = 0,
limit: int = 50,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> list[MessageReplyToResponse]:
with get_db_context(db) as db:
all_messages = (
db.query(Message)
async with get_async_db_context(db) as db:
result = await db.execute(
select(Message)
.filter_by(channel_id=channel_id, parent_id=None)
.order_by(Message.created_at.desc())
.offset(skip)
.limit(limit)
.all()
)
all_messages = result.scalars().all()
messages = []
for message in all_messages:
reply_to_message = (
self.get_message_by_id(message.reply_to_id, include_thread_replies=False, db=db)
await self.get_message_by_id(message.reply_to_id, include_thread_replies=False, db=db)
if message.reply_to_id
else None
)
webhook_info = message.meta.get('webhook') if message.meta else None
user_info = None
if webhook_info and webhook_info.get('id'):
webhook = Channels.get_webhook_by_id(webhook_info.get('id'), db=db)
if webhook:
user_info = {
'id': webhook.id,
'name': webhook.name,
'role': 'webhook',
}
else:
user_info = {
'id': webhook_info.get('id'),
'name': 'Deleted Webhook',
'role': 'webhook',
}
user_info = await self._resolve_user_info(message, db)
messages.append(
MessageReplyToResponse.model_validate(
@ -327,28 +321,28 @@ class MessageTable:
)
return messages
def get_messages_by_parent_id(
async def get_messages_by_parent_id(
self,
channel_id: str,
parent_id: str,
skip: int = 0,
limit: int = 50,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> list[MessageReplyToResponse]:
with get_db_context(db) as db:
message = db.get(Message, parent_id)
async with get_async_db_context(db) as db:
message = await db.get(Message, parent_id)
if not message:
return []
all_messages = (
db.query(Message)
result = await db.execute(
select(Message)
.filter_by(channel_id=channel_id, parent_id=parent_id)
.order_by(Message.created_at.desc())
.offset(skip)
.limit(limit)
.all()
)
all_messages = list(result.scalars().all())
# If length of all_messages is less than limit, then add the parent message
if len(all_messages) < limit:
@ -357,27 +351,12 @@ class MessageTable:
messages = []
for message in all_messages:
reply_to_message = (
self.get_message_by_id(message.reply_to_id, include_thread_replies=False, db=db)
await self.get_message_by_id(message.reply_to_id, include_thread_replies=False, db=db)
if message.reply_to_id
else None
)
webhook_info = message.meta.get('webhook') if message.meta else None
user_info = None
if webhook_info and webhook_info.get('id'):
webhook = Channels.get_webhook_by_id(webhook_info.get('id'), db=db)
if webhook:
user_info = {
'id': webhook.id,
'name': webhook.name,
'role': 'webhook',
}
else:
user_info = {
'id': webhook_info.get('id'),
'name': 'Deleted Webhook',
'role': 'webhook',
}
user_info = await self._resolve_user_info(message, db)
messages.append(
MessageReplyToResponse.model_validate(
@ -390,34 +369,37 @@ class MessageTable:
)
return messages
def get_last_message_by_channel_id(self, channel_id: str, db: Optional[Session] = None) -> Optional[MessageModel]:
with get_db_context(db) as db:
message = db.query(Message).filter_by(channel_id=channel_id).order_by(Message.created_at.desc()).first()
async def get_last_message_by_channel_id(self, channel_id: str, db: Optional[AsyncSession] = None) -> Optional[MessageModel]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(Message).filter_by(channel_id=channel_id).order_by(Message.created_at.desc()).limit(1)
)
message = result.scalars().first()
return MessageModel.model_validate(message) if message else None
def get_pinned_messages_by_channel_id(
async def get_pinned_messages_by_channel_id(
self,
channel_id: str,
skip: int = 0,
limit: int = 50,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> list[MessageModel]:
with get_db_context(db) as db:
all_messages = (
db.query(Message)
async with get_async_db_context(db) as db:
result = await db.execute(
select(Message)
.filter_by(channel_id=channel_id, is_pinned=True)
.order_by(Message.pinned_at.desc())
.offset(skip)
.limit(limit)
.all()
)
all_messages = result.scalars().all()
return [MessageModel.model_validate(message) for message in all_messages]
def update_message_by_id(
self, id: str, form_data: MessageForm, db: Optional[Session] = None
async def update_message_by_id(
self, id: str, form_data: MessageForm, db: Optional[AsyncSession] = None
) -> Optional[MessageModel]:
with get_db_context(db) as db:
message = db.get(Message, id)
async with get_async_db_context(db) as db:
message = await db.get(Message, id)
message.content = form_data.content
message.data = {
**(message.data if message.data else {}),
@ -428,49 +410,53 @@ class MessageTable:
**(form_data.meta if form_data.meta else {}),
}
message.updated_at = int(time.time_ns())
db.commit()
db.refresh(message)
await db.commit()
await db.refresh(message)
return MessageModel.model_validate(message) if message else None
def update_is_pinned_by_id(
async def update_is_pinned_by_id(
self,
id: str,
is_pinned: bool,
pinned_by: Optional[str] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[MessageModel]:
with get_db_context(db) as db:
message = db.get(Message, id)
async with get_async_db_context(db) as db:
message = await db.get(Message, id)
message.is_pinned = is_pinned
message.pinned_at = int(time.time_ns()) if is_pinned else None
message.pinned_by = pinned_by if is_pinned else None
db.commit()
db.refresh(message)
await db.commit()
await db.refresh(message)
return MessageModel.model_validate(message) if message else None
def get_unread_message_count(
async def get_unread_message_count(
self,
channel_id: str,
user_id: str,
last_read_at: Optional[int] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> int:
with get_db_context(db) as db:
query = db.query(Message).filter(
async with get_async_db_context(db) as db:
stmt = select(func.count(Message.id)).filter(
Message.channel_id == channel_id,
Message.parent_id == None, # only count top-level messages
Message.created_at > (last_read_at if last_read_at else 0),
)
if user_id:
query = query.filter(Message.user_id != user_id)
return query.count()
stmt = stmt.filter(Message.user_id != user_id)
result = await db.execute(stmt)
return result.scalar()
def add_reaction_to_message(
self, id: str, user_id: str, name: str, db: Optional[Session] = None
async def add_reaction_to_message(
self, id: str, user_id: str, name: str, db: Optional[AsyncSession] = None
) -> Optional[MessageReactionModel]:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
# check for existing reaction
existing_reaction = db.query(MessageReaction).filter_by(message_id=id, user_id=user_id, name=name).first()
result = await db.execute(
select(MessageReaction).filter_by(message_id=id, user_id=user_id, name=name)
)
existing_reaction = result.scalars().first()
if existing_reaction:
return MessageReactionModel.model_validate(existing_reaction)
@ -484,19 +470,19 @@ class MessageTable:
)
result = MessageReaction(**reaction.model_dump())
db.add(result)
db.commit()
db.refresh(result)
await db.commit()
await db.refresh(result)
return MessageReactionModel.model_validate(result) if result else None
def get_reactions_by_message_id(self, id: str, db: Optional[Session] = None) -> list[Reactions]:
with get_db_context(db) as db:
async def get_reactions_by_message_id(self, id: str, db: Optional[AsyncSession] = None) -> list[Reactions]:
async with get_async_db_context(db) as db:
# JOIN User so all user info is fetched in one query
results = (
db.query(MessageReaction, User)
result = await db.execute(
select(MessageReaction, User)
.join(User, MessageReaction.user_id == User.id)
.filter(MessageReaction.message_id == id)
.all()
)
results = result.all()
reactions = {}
@ -518,58 +504,60 @@ class MessageTable:
return [Reactions(**reaction) for reaction in reactions.values()]
def remove_reaction_by_id_and_user_id_and_name(
self, id: str, user_id: str, name: str, db: Optional[Session] = None
async def remove_reaction_by_id_and_user_id_and_name(
self, id: str, user_id: str, name: str, db: Optional[AsyncSession] = None
) -> bool:
with get_db_context(db) as db:
db.query(MessageReaction).filter_by(message_id=id, user_id=user_id, name=name).delete()
db.commit()
async with get_async_db_context(db) as db:
await db.execute(delete(MessageReaction).filter_by(message_id=id, user_id=user_id, name=name))
await db.commit()
return True
def delete_reactions_by_id(self, id: str, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
db.query(MessageReaction).filter_by(message_id=id).delete()
db.commit()
async def delete_reactions_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
await db.execute(delete(MessageReaction).filter_by(message_id=id))
await db.commit()
return True
def delete_replies_by_id(self, id: str, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
db.query(Message).filter_by(parent_id=id).delete()
db.commit()
async def delete_replies_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
await db.execute(delete(Message).filter_by(parent_id=id))
await db.commit()
return True
def delete_message_by_id(self, id: str, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
db.query(Message).filter_by(id=id).delete()
async def delete_message_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
await db.execute(delete(Message).filter_by(id=id))
# Delete all reactions to this message
db.query(MessageReaction).filter_by(message_id=id).delete()
await db.execute(delete(MessageReaction).filter_by(message_id=id))
db.commit()
await db.commit()
return True
def search_messages_by_channel_ids(
async def search_messages_by_channel_ids(
self,
channel_ids: list[str],
query: str,
start_timestamp: Optional[int] = None,
end_timestamp: Optional[int] = None,
limit: int = 10,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> list[MessageModel]:
"""Search messages in specified channels by content."""
with get_db_context(db) as db:
query_builder = db.query(Message).filter(
async with get_async_db_context(db) as db:
stmt = select(Message).filter(
Message.channel_id.in_(channel_ids),
Message.content.ilike(f'%{query}%'),
)
if start_timestamp:
query_builder = query_builder.filter(Message.created_at >= start_timestamp)
stmt = stmt.filter(Message.created_at >= start_timestamp)
if end_timestamp:
query_builder = query_builder.filter(Message.created_at <= end_timestamp)
stmt = stmt.filter(Message.created_at <= end_timestamp)
messages = query_builder.order_by(Message.created_at.desc()).limit(limit).all()
stmt = stmt.order_by(Message.created_at.desc()).limit(limit)
result = await db.execute(stmt)
messages = result.scalars().all()
return [MessageModel.model_validate(msg) for msg in messages]

View file

@ -2,8 +2,9 @@ import logging
import time
from typing import Optional
from sqlalchemy.orm import Session
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
from sqlalchemy import select, delete, update, or_, func, String, cast
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import Base, JSONField, get_async_db_context
from open_webui.models.groups import Groups
from open_webui.models.users import User, UserModel, Users, UserResponse
@ -12,9 +13,6 @@ from open_webui.models.access_grants import AccessGrantModel, AccessGrants
from pydantic import BaseModel, ConfigDict, Field, model_validator
from sqlalchemy import String, cast, or_, and_, func
from sqlalchemy.dialects import postgresql, sqlite
from sqlalchemy.dialects.postgresql import JSONB
from sqlalchemy import BigInteger, Column, Text, Boolean
@ -154,26 +152,26 @@ class ModelForm(BaseModel):
class ModelsTable:
def _get_access_grants(self, model_id: str, db: Optional[Session] = None) -> list[AccessGrantModel]:
return AccessGrants.get_grants_by_resource('model', model_id, db=db)
async def _get_access_grants(self, model_id: str, db: Optional[AsyncSession] = None) -> list[AccessGrantModel]:
return await AccessGrants.get_grants_by_resource('model', model_id, db=db)
def _to_model_model(
async def _to_model_model(
self,
model: Model,
access_grants: Optional[list[AccessGrantModel]] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> ModelModel:
model_data = ModelModel.model_validate(model).model_dump(exclude={'access_grants'})
model_data['access_grants'] = (
access_grants if access_grants is not None else self._get_access_grants(model_data['id'], db=db)
access_grants if access_grants is not None else await self._get_access_grants(model_data['id'], db=db)
)
return ModelModel.model_validate(model_data)
def insert_new_model(
self, form_data: ModelForm, user_id: str, db: Optional[Session] = None
async def insert_new_model(
self, form_data: ModelForm, user_id: str, db: Optional[AsyncSession] = None
) -> Optional[ModelModel]:
try:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
result = Model(
**{
**form_data.model_dump(exclude={'access_grants'}),
@ -183,37 +181,39 @@ class ModelsTable:
}
)
db.add(result)
db.commit()
db.refresh(result)
AccessGrants.set_access_grants('model', result.id, form_data.access_grants, db=db)
await db.commit()
await db.refresh(result)
await AccessGrants.set_access_grants('model', result.id, form_data.access_grants, db=db)
if result:
return self._to_model_model(result, db=db)
return await self._to_model_model(result, db=db)
else:
return None
except Exception as e:
log.exception(f'Failed to insert a new model: {e}')
return None
def get_all_models(self, db: Optional[Session] = None) -> list[ModelModel]:
with get_db_context(db) as db:
all_models = db.query(Model).all()
async def get_all_models(self, db: Optional[AsyncSession] = None) -> list[ModelModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(Model))
all_models = result.scalars().all()
model_ids = [model.id for model in all_models]
grants_map = AccessGrants.get_grants_by_resources('model', model_ids, db=db)
grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db)
return [
self._to_model_model(model, access_grants=grants_map.get(model.id, []), db=db) for model in all_models
await self._to_model_model(model, access_grants=grants_map.get(model.id, []), db=db) for model in all_models
]
def get_models(self, db: Optional[Session] = None) -> list[ModelUserResponse]:
with get_db_context(db) as db:
all_models = db.query(Model).filter(Model.base_model_id != None).all()
async def get_models(self, db: Optional[AsyncSession] = None) -> list[ModelUserResponse]:
async with get_async_db_context(db) as db:
result = await db.execute(select(Model).filter(Model.base_model_id != None))
all_models = result.scalars().all()
user_ids = list(set(model.user_id for model in all_models))
model_ids = [model.id for model in all_models]
users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
users = await Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
users_dict = {user.id: user for user in users}
grants_map = AccessGrants.get_grants_by_resources('model', model_ids, db=db)
grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db)
models = []
for model in all_models:
@ -221,44 +221,48 @@ class ModelsTable:
models.append(
ModelUserResponse.model_validate(
{
**self._to_model_model(
**(await self._to_model_model(
model,
access_grants=grants_map.get(model.id, []),
db=db,
).model_dump(),
)).model_dump(),
'user': user.model_dump() if user else None,
}
)
)
return models
def get_base_models(self, db: Optional[Session] = None) -> list[ModelModel]:
with get_db_context(db) as db:
all_models = db.query(Model).filter(Model.base_model_id == None).all()
async def get_base_models(self, db: Optional[AsyncSession] = None) -> list[ModelModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(Model).filter(Model.base_model_id == None))
all_models = result.scalars().all()
model_ids = [model.id for model in all_models]
grants_map = AccessGrants.get_grants_by_resources('model', model_ids, db=db)
grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db)
return [
self._to_model_model(model, access_grants=grants_map.get(model.id, []), db=db) for model in all_models
await self._to_model_model(model, access_grants=grants_map.get(model.id, []), db=db) for model in all_models
]
def get_models_by_user_id(
self, user_id: str, permission: str = 'write', db: Optional[Session] = None
async def get_models_by_user_id(
self, user_id: str, permission: str = 'write', db: Optional[AsyncSession] = None
) -> list[ModelUserResponse]:
models = self.get_models(db=db)
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id, db=db)}
return [
model
for model in models
if model.user_id == user_id
or AccessGrants.has_access(
models = await self.get_models(db=db)
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
user_group_ids = {group.id for group in user_groups}
result = []
for model in models:
if model.user_id == user_id:
result.append(model)
elif await AccessGrants.has_access(
user_id=user_id,
resource_type='model',
resource_id=model.id,
permission=permission,
user_group_ids=user_group_ids,
db=db,
)
]
):
result.append(model)
return result
def _has_permission(self, db, query, filter: dict, permission: str = 'read'):
return AccessGrants.has_permission_filter(
@ -270,23 +274,22 @@ class ModelsTable:
permission=permission,
)
def search_models(
async def search_models(
self,
user_id: str,
filter: dict = {},
skip: int = 0,
limit: int = 30,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> ModelListResponse:
with get_db_context(db) as db:
# Join GroupMember so we can order by group_id when requested
query = db.query(Model, User).outerjoin(User, User.id == Model.user_id)
query = query.filter(Model.base_model_id != None)
async with get_async_db_context(db) as db:
stmt = select(Model, User).outerjoin(User, User.id == Model.user_id)
stmt = stmt.filter(Model.base_model_id != None)
if filter:
query_key = filter.get('query')
if query_key:
query = query.filter(
stmt = stmt.filter(
or_(
Model.name.ilike(f'%{query_key}%'),
Model.base_model_id.ilike(f'%{query_key}%'),
@ -298,92 +301,95 @@ class ModelsTable:
view_option = filter.get('view_option')
if view_option == 'created':
query = query.filter(Model.user_id == user_id)
stmt = stmt.filter(Model.user_id == user_id)
elif view_option == 'shared':
query = query.filter(Model.user_id != user_id)
stmt = stmt.filter(Model.user_id != user_id)
# Apply access control filtering
query = self._has_permission(
stmt = self._has_permission(
db,
query,
stmt,
filter,
permission='read',
)
tag = filter.get('tag')
if tag:
# TODO: This is a simple implementation and should be improved for performance
like_pattern = f'%"{tag.lower()}"%' # `"tag"` inside JSON array
like_pattern = f'%"{tag.lower()}"%'
meta_text = func.lower(cast(Model.meta, String))
query = query.filter(meta_text.like(like_pattern))
stmt = stmt.filter(meta_text.like(like_pattern))
order_by = filter.get('order_by')
direction = filter.get('direction')
if order_by == 'name':
if direction == 'asc':
query = query.order_by(Model.name.asc())
stmt = stmt.order_by(Model.name.asc())
else:
query = query.order_by(Model.name.desc())
stmt = stmt.order_by(Model.name.desc())
elif order_by == 'created_at':
if direction == 'asc':
query = query.order_by(Model.created_at.asc())
stmt = stmt.order_by(Model.created_at.asc())
else:
query = query.order_by(Model.created_at.desc())
stmt = stmt.order_by(Model.created_at.desc())
elif order_by == 'updated_at':
if direction == 'asc':
query = query.order_by(Model.updated_at.asc())
stmt = stmt.order_by(Model.updated_at.asc())
else:
query = query.order_by(Model.updated_at.desc())
stmt = stmt.order_by(Model.updated_at.desc())
else:
query = query.order_by(Model.created_at.desc())
stmt = stmt.order_by(Model.created_at.desc())
# Count BEFORE pagination
total = query.count()
count_result = await db.execute(
select(func.count()).select_from(stmt.subquery())
)
total = count_result.scalar()
if skip:
query = query.offset(skip)
stmt = stmt.offset(skip)
if limit:
query = query.limit(limit)
stmt = stmt.limit(limit)
items = query.all()
result = await db.execute(stmt)
items = result.all()
model_ids = [model.id for model, _ in items]
grants_map = AccessGrants.get_grants_by_resources('model', model_ids, db=db)
grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db)
models = []
for model, user in items:
models.append(
ModelUserResponse(
**self._to_model_model(
**(await self._to_model_model(
model,
access_grants=grants_map.get(model.id, []),
db=db,
).model_dump(),
)).model_dump(),
user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
)
)
return ModelListResponse(items=models, total=total)
def get_model_by_id(self, id: str, db: Optional[Session] = None) -> Optional[ModelModel]:
async def get_model_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[ModelModel]:
try:
with get_db_context(db) as db:
model = db.get(Model, id)
return self._to_model_model(model, db=db) if model else None
async with get_async_db_context(db) as db:
model = await db.get(Model, id)
return await self._to_model_model(model, db=db) if model else None
except Exception:
return None
def get_models_by_ids(self, ids: list[str], db: Optional[Session] = None) -> list[ModelModel]:
async def get_models_by_ids(self, ids: list[str], db: Optional[AsyncSession] = None) -> list[ModelModel]:
try:
with get_db_context(db) as db:
models = db.query(Model).filter(Model.id.in_(ids)).all()
async with get_async_db_context(db) as db:
result = await db.execute(select(Model).filter(Model.id.in_(ids)))
models = result.scalars().all()
model_ids = [model.id for model in models]
grants_map = AccessGrants.get_grants_by_resources('model', model_ids, db=db)
grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db)
return [
self._to_model_model(
await self._to_model_model(
model,
access_grants=grants_map.get(model.id, []),
db=db,
@ -393,82 +399,86 @@ class ModelsTable:
except Exception:
return []
def toggle_model_by_id(self, id: str, db: Optional[Session] = None) -> Optional[ModelModel]:
with get_db_context(db) as db:
async def toggle_model_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[ModelModel]:
async with get_async_db_context(db) as db:
try:
model = db.query(Model).filter_by(id=id).first()
result = await db.execute(select(Model).filter_by(id=id))
model = result.scalars().first()
if not model:
return None
model.is_active = not model.is_active
model.updated_at = int(time.time())
db.commit()
db.refresh(model)
await db.commit()
await db.refresh(model)
return self._to_model_model(model, db=db)
return await self._to_model_model(model, db=db)
except Exception:
return None
def update_model_by_id(self, id: str, model: ModelForm, db: Optional[Session] = None) -> Optional[ModelModel]:
async def update_model_by_id(self, id: str, model: ModelForm, db: Optional[AsyncSession] = None) -> Optional[ModelModel]:
try:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
# update only the fields that are present in the model
data = model.model_dump(exclude={'id', 'access_grants'})
data['updated_at'] = int(time.time())
result = db.query(Model).filter_by(id=id).update(data)
await db.execute(update(Model).filter_by(id=id).values(**data))
db.commit()
await db.commit()
if model.access_grants is not None:
AccessGrants.set_access_grants('model', id, model.access_grants, db=db)
await AccessGrants.set_access_grants('model', id, model.access_grants, db=db)
return self.get_model_by_id(id, db=db)
return await self.get_model_by_id(id, db=db)
except Exception as e:
log.exception(f'Failed to update the model by id {id}: {e}')
return None
def update_model_updated_at_by_id(self, id: str, db: Optional[Session] = None) -> Optional[ModelModel]:
async def update_model_updated_at_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[ModelModel]:
try:
with get_db_context(db) as db:
result = db.query(Model).filter_by(id=id).first()
if not result:
async with get_async_db_context(db) as db:
result = await db.execute(select(Model).filter_by(id=id))
model_obj = result.scalars().first()
if not model_obj:
return None
result.updated_at = int(time.time())
db.commit()
db.refresh(result)
return self._to_model_model(result, db=db)
model_obj.updated_at = int(time.time())
await db.commit()
await db.refresh(model_obj)
return await self._to_model_model(model_obj, db=db)
except Exception as e:
log.exception(f'Failed to update the model updated_at by id {id}: {e}')
return None
def delete_model_by_id(self, id: str, db: Optional[Session] = None) -> bool:
async def delete_model_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
try:
with get_db_context(db) as db:
AccessGrants.revoke_all_access('model', id, db=db)
db.query(Model).filter_by(id=id).delete()
db.commit()
async with get_async_db_context(db) as db:
await AccessGrants.revoke_all_access('model', id, db=db)
await db.execute(delete(Model).filter_by(id=id))
await db.commit()
return True
except Exception:
return False
def delete_all_models(self, db: Optional[Session] = None) -> bool:
async def delete_all_models(self, db: Optional[AsyncSession] = None) -> bool:
try:
with get_db_context(db) as db:
model_ids = [row[0] for row in db.query(Model.id).all()]
async with get_async_db_context(db) as db:
result = await db.execute(select(Model.id))
model_ids = [row[0] for row in result.all()]
for model_id in model_ids:
AccessGrants.revoke_all_access('model', model_id, db=db)
db.query(Model).delete()
db.commit()
await AccessGrants.revoke_all_access('model', model_id, db=db)
await db.execute(delete(Model))
await db.commit()
return True
except Exception:
return False
def sync_models(self, user_id: str, models: list[ModelModel], db: Optional[Session] = None) -> list[ModelModel]:
async def sync_models(self, user_id: str, models: list[ModelModel], db: Optional[AsyncSession] = None) -> list[ModelModel]:
try:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
# Get existing models
existing_models = db.query(Model).all()
result = await db.execute(select(Model))
existing_models = result.scalars().all()
existing_ids = {model.id for model in existing_models}
# Prepare a set of new model IDs
@ -477,12 +487,12 @@ class ModelsTable:
# Update or insert models
for model in models:
if model.id in existing_ids:
db.query(Model).filter_by(id=model.id).update(
{
await db.execute(
update(Model).filter_by(id=model.id).values(
**model.model_dump(exclude={'access_grants'}),
'user_id': user_id,
'updated_at': int(time.time()),
}
user_id=user_id,
updated_at=int(time.time()),
)
)
else:
new_model = Model(
@ -493,21 +503,22 @@ class ModelsTable:
}
)
db.add(new_model)
AccessGrants.set_access_grants('model', model.id, model.access_grants, db=db)
await AccessGrants.set_access_grants('model', model.id, model.access_grants, db=db)
# Remove models that are no longer present
for model in existing_models:
if model.id not in new_model_ids:
AccessGrants.revoke_all_access('model', model.id, db=db)
db.delete(model)
await AccessGrants.revoke_all_access('model', model.id, db=db)
await db.delete(model)
db.commit()
await db.commit()
all_models = db.query(Model).all()
result = await db.execute(select(Model))
all_models = result.scalars().all()
model_ids = [model.id for model in all_models]
grants_map = AccessGrants.get_grants_by_resources('model', model_ids, db=db)
grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db)
return [
self._to_model_model(
await self._to_model_model(
model,
access_grants=grants_map.get(model.id, []),
db=db,

View file

@ -4,8 +4,9 @@ import uuid
from typing import Optional
from functools import lru_cache
from sqlalchemy.orm import Session
from open_webui.internal.db import Base, get_db, get_db_context
from sqlalchemy import select, delete, update, or_, func, cast
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import Base, get_async_db_context
from open_webui.models.groups import Groups
from open_webui.models.users import User, UserModel, Users, UserResponse
from open_webui.models.access_grants import AccessGrantModel, AccessGrants
@ -13,7 +14,6 @@ from open_webui.models.access_grants import AccessGrantModel, AccessGrants
from pydantic import BaseModel, ConfigDict, Field
from sqlalchemy import BigInteger, Column, Text, JSON
from sqlalchemy import or_, func, cast
####################
# Note DB Schema
@ -88,18 +88,18 @@ class NoteListResponse(BaseModel):
class NoteTable:
def _get_access_grants(self, note_id: str, db: Optional[Session] = None) -> list[AccessGrantModel]:
return AccessGrants.get_grants_by_resource('note', note_id, db=db)
async def _get_access_grants(self, note_id: str, db: Optional[AsyncSession] = None) -> list[AccessGrantModel]:
return await AccessGrants.get_grants_by_resource('note', note_id, db=db)
def _to_note_model(
async def _to_note_model(
self,
note: Note,
access_grants: Optional[list[AccessGrantModel]] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> NoteModel:
note_data = NoteModel.model_validate(note).model_dump(exclude={'access_grants'})
note_data['access_grants'] = (
access_grants if access_grants is not None else self._get_access_grants(note_data['id'], db=db)
access_grants if access_grants is not None else await self._get_access_grants(note_data['id'], db=db)
)
return NoteModel.model_validate(note_data)
@ -113,8 +113,8 @@ class NoteTable:
permission=permission,
)
def insert_new_note(self, user_id: str, form_data: NoteForm, db: Optional[Session] = None) -> Optional[NoteModel]:
with get_db_context(db) as db:
async def insert_new_note(self, user_id: str, form_data: NoteForm, db: Optional[AsyncSession] = None) -> Optional[NoteModel]:
async with get_async_db_context(db) as db:
note = NoteModel(
**{
'id': str(uuid.uuid4()),
@ -129,38 +129,39 @@ class NoteTable:
new_note = Note(**note.model_dump(exclude={'access_grants'}))
db.add(new_note)
db.commit()
AccessGrants.set_access_grants('note', note.id, form_data.access_grants, db=db)
return self._to_note_model(new_note, db=db)
await db.commit()
await AccessGrants.set_access_grants('note', note.id, form_data.access_grants, db=db)
return await self._to_note_model(new_note, db=db)
def get_notes(self, skip: int = 0, limit: int = 50, db: Optional[Session] = None) -> list[NoteModel]:
with get_db_context(db) as db:
query = db.query(Note).order_by(Note.updated_at.desc())
async def get_notes(self, skip: int = 0, limit: int = 50, db: Optional[AsyncSession] = None) -> list[NoteModel]:
async with get_async_db_context(db) as db:
stmt = select(Note).order_by(Note.updated_at.desc())
if skip is not None:
query = query.offset(skip)
stmt = stmt.offset(skip)
if limit is not None:
query = query.limit(limit)
notes = query.all()
stmt = stmt.limit(limit)
result = await db.execute(stmt)
notes = result.scalars().all()
note_ids = [note.id for note in notes]
grants_map = AccessGrants.get_grants_by_resources('note', note_ids, db=db)
return [self._to_note_model(note, access_grants=grants_map.get(note.id, []), db=db) for note in notes]
grants_map = await AccessGrants.get_grants_by_resources('note', note_ids, db=db)
return [await self._to_note_model(note, access_grants=grants_map.get(note.id, []), db=db) for note in notes]
def search_notes(
async def search_notes(
self,
user_id: str,
filter: dict = {},
skip: int = 0,
limit: int = 30,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> NoteListResponse:
with get_db_context(db) as db:
query = db.query(Note, User).outerjoin(User, User.id == Note.user_id)
async with get_async_db_context(db) as db:
stmt = select(Note, User).outerjoin(User, User.id == Note.user_id)
if filter:
query_key = filter.get('query')
if query_key:
# Normalize search by removing hyphens and spaces (e.g., "todo" matches "to-do" and "to do")
normalized_query = query_key.replace('-', '').replace(' ', '')
query = query.filter(
stmt = stmt.filter(
or_(
func.replace(func.replace(Note.title, '-', ''), ' ', '').ilike(f'%{normalized_query}%'),
func.replace(
@ -173,9 +174,9 @@ class NoteTable:
view_option = filter.get('view_option')
if view_option == 'created':
query = query.filter(Note.user_id == user_id)
stmt = stmt.filter(Note.user_id == user_id)
elif view_option == 'shared':
query = query.filter(Note.user_id != user_id)
stmt = stmt.filter(Note.user_id != user_id)
# Apply access control filtering
if 'permission' in filter:
@ -183,9 +184,9 @@ class NoteTable:
else:
permission = 'write'
query = self._has_permission(
stmt = self._has_permission(
db,
query,
stmt,
filter,
permission=permission,
)
@ -195,87 +196,95 @@ class NoteTable:
if order_by == 'name':
if direction == 'asc':
query = query.order_by(Note.title.asc())
stmt = stmt.order_by(Note.title.asc())
else:
query = query.order_by(Note.title.desc())
stmt = stmt.order_by(Note.title.desc())
elif order_by == 'created_at':
if direction == 'asc':
query = query.order_by(Note.created_at.asc())
stmt = stmt.order_by(Note.created_at.asc())
else:
query = query.order_by(Note.created_at.desc())
stmt = stmt.order_by(Note.created_at.desc())
elif order_by == 'updated_at':
if direction == 'asc':
query = query.order_by(Note.updated_at.asc())
stmt = stmt.order_by(Note.updated_at.asc())
else:
query = query.order_by(Note.updated_at.desc())
stmt = stmt.order_by(Note.updated_at.desc())
else:
query = query.order_by(Note.updated_at.desc())
stmt = stmt.order_by(Note.updated_at.desc())
else:
query = query.order_by(Note.updated_at.desc())
stmt = stmt.order_by(Note.updated_at.desc())
# Count BEFORE pagination
total = query.count()
count_result = await db.execute(
select(func.count()).select_from(stmt.subquery())
)
total = count_result.scalar()
if skip:
query = query.offset(skip)
stmt = stmt.offset(skip)
if limit:
query = query.limit(limit)
stmt = stmt.limit(limit)
items = query.all()
result = await db.execute(stmt)
items = result.all()
note_ids = [note.id for note, _ in items]
grants_map = AccessGrants.get_grants_by_resources('note', note_ids, db=db)
grants_map = await AccessGrants.get_grants_by_resources('note', note_ids, db=db)
notes = []
for note, user in items:
notes.append(
NoteUserResponse(
**self._to_note_model(
**(await self._to_note_model(
note,
access_grants=grants_map.get(note.id, []),
db=db,
).model_dump(),
)).model_dump(),
user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
)
)
return NoteListResponse(items=notes, total=total)
def get_notes_by_user_id(
async def get_notes_by_user_id(
self,
user_id: str,
permission: str = 'read',
skip: int = 0,
limit: int = 50,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> list[NoteModel]:
with get_db_context(db) as db:
user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id, db=db)]
async with get_async_db_context(db) as db:
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
user_group_ids = [group.id for group in user_groups]
query = db.query(Note).order_by(Note.updated_at.desc())
query = self._has_permission(db, query, {'user_id': user_id, 'group_ids': user_group_ids}, permission)
stmt = select(Note).order_by(Note.updated_at.desc())
stmt = self._has_permission(db, stmt, {'user_id': user_id, 'group_ids': user_group_ids}, permission)
if skip is not None:
query = query.offset(skip)
stmt = stmt.offset(skip)
if limit is not None:
query = query.limit(limit)
stmt = stmt.limit(limit)
notes = query.all()
result = await db.execute(stmt)
notes = result.scalars().all()
note_ids = [note.id for note in notes]
grants_map = AccessGrants.get_grants_by_resources('note', note_ids, db=db)
return [self._to_note_model(note, access_grants=grants_map.get(note.id, []), db=db) for note in notes]
grants_map = await AccessGrants.get_grants_by_resources('note', note_ids, db=db)
return [await self._to_note_model(note, access_grants=grants_map.get(note.id, []), db=db) for note in notes]
def get_note_by_id(self, id: str, db: Optional[Session] = None) -> Optional[NoteModel]:
with get_db_context(db) as db:
note = db.query(Note).filter(Note.id == id).first()
return self._to_note_model(note, db=db) if note else None
async def get_note_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[NoteModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(Note).filter(Note.id == id))
note = result.scalars().first()
return await self._to_note_model(note, db=db) if note else None
def update_note_by_id(
self, id: str, form_data: NoteUpdateForm, db: Optional[Session] = None
async def update_note_by_id(
self, id: str, form_data: NoteUpdateForm, db: Optional[AsyncSession] = None
) -> Optional[NoteModel]:
with get_db_context(db) as db:
note = db.query(Note).filter(Note.id == id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(Note).filter(Note.id == id))
note = result.scalars().first()
if not note:
return None
@ -289,19 +298,19 @@ class NoteTable:
note.meta = {**note.meta, **form_data['meta']}
if 'access_grants' in form_data:
AccessGrants.set_access_grants('note', id, form_data['access_grants'], db=db)
await AccessGrants.set_access_grants('note', id, form_data['access_grants'], db=db)
note.updated_at = int(time.time_ns())
db.commit()
return self._to_note_model(note, db=db) if note else None
await db.commit()
return await self._to_note_model(note, db=db) if note else None
def delete_note_by_id(self, id: str, db: Optional[Session] = None) -> bool:
async def delete_note_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
try:
with get_db_context(db) as db:
AccessGrants.revoke_all_access('note', id, db=db)
db.query(Note).filter(Note.id == id).delete()
db.commit()
async with get_async_db_context(db) as db:
await AccessGrants.revoke_all_access('note', id, db=db)
await db.execute(delete(Note).filter(Note.id == id))
await db.commit()
return True
except Exception:
return False

View file

@ -8,8 +8,9 @@ import json
from cryptography.fernet import Fernet
from sqlalchemy.orm import Session
from open_webui.internal.db import Base, get_db, get_db_context
from sqlalchemy import select, delete, update
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import Base, get_async_db_context
from open_webui.env import OAUTH_SESSION_TOKEN_ENCRYPTION_KEY
from pydantic import BaseModel, ConfigDict
@ -103,16 +104,16 @@ class OAuthSessionTable:
log.error(f'Error decrypting tokens: {type(e).__name__}: {e}')
raise
def create_session(
async def create_session(
self,
user_id: str,
provider: str,
token: dict,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[OAuthSessionModel]:
"""Create a new OAuth session"""
try:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
current_time = int(time.time())
id = str(uuid.uuid4())
@ -129,91 +130,126 @@ class OAuthSessionTable:
)
db.add(result)
db.commit()
db.refresh(result)
await db.commit()
await db.refresh(result)
if result:
db.expunge(result) # Detach so dict swap is never flushed
result.token = token # Return decrypted token
return OAuthSessionModel.model_validate(result)
# Make a copy of the model data before closing session
model = OAuthSessionModel(
id=result.id,
user_id=result.user_id,
provider=result.provider,
token=token, # Return decrypted token
expires_at=result.expires_at,
created_at=result.created_at,
updated_at=result.updated_at,
)
return model
else:
return None
except Exception as e:
log.error(f'Error creating OAuth session: {e}')
return None
def get_session_by_id(self, session_id: str, db: Optional[Session] = None) -> Optional[OAuthSessionModel]:
async def get_session_by_id(self, session_id: str, db: Optional[AsyncSession] = None) -> Optional[OAuthSessionModel]:
"""Get OAuth session by ID"""
try:
with get_db_context(db) as db:
session = db.query(OAuthSession).filter_by(id=session_id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(OAuthSession).filter_by(id=session_id))
session = result.scalars().first()
if session:
db.expunge(session)
session.token = self._decrypt_token(session.token)
return OAuthSessionModel.model_validate(session)
return OAuthSessionModel(
id=session.id,
user_id=session.user_id,
provider=session.provider,
token=self._decrypt_token(session.token),
expires_at=session.expires_at,
created_at=session.created_at,
updated_at=session.updated_at,
)
return None
except Exception as e:
log.error(f'Error getting OAuth session by ID: {e}')
return None
def get_session_by_id_and_user_id(
self, session_id: str, user_id: str, db: Optional[Session] = None
async def get_session_by_id_and_user_id(
self, session_id: str, user_id: str, db: Optional[AsyncSession] = None
) -> Optional[OAuthSessionModel]:
"""Get OAuth session by ID and user ID"""
try:
with get_db_context(db) as db:
session = db.query(OAuthSession).filter_by(id=session_id, user_id=user_id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(OAuthSession).filter_by(id=session_id, user_id=user_id))
session = result.scalars().first()
if session:
db.expunge(session)
session.token = self._decrypt_token(session.token)
return OAuthSessionModel.model_validate(session)
return OAuthSessionModel(
id=session.id,
user_id=session.user_id,
provider=session.provider,
token=self._decrypt_token(session.token),
expires_at=session.expires_at,
created_at=session.created_at,
updated_at=session.updated_at,
)
return None
except Exception as e:
log.error(f'Error getting OAuth session by ID: {e}')
return None
def get_session_by_provider_and_user_id(
self, provider: str, user_id: str, db: Optional[Session] = None
async def get_session_by_provider_and_user_id(
self, provider: str, user_id: str, db: Optional[AsyncSession] = None
) -> Optional[OAuthSessionModel]:
"""Get OAuth session by provider and user ID"""
try:
with get_db_context(db) as db:
session = (
db.query(OAuthSession)
async with get_async_db_context(db) as db:
result = await db.execute(
select(OAuthSession)
.filter_by(provider=provider, user_id=user_id)
.order_by(OAuthSession.created_at.desc())
.first()
)
session = result.scalars().first()
if session:
db.expunge(session)
session.token = self._decrypt_token(session.token)
return OAuthSessionModel.model_validate(session)
return OAuthSessionModel(
id=session.id,
user_id=session.user_id,
provider=session.provider,
token=self._decrypt_token(session.token),
expires_at=session.expires_at,
created_at=session.created_at,
updated_at=session.updated_at,
)
return None
except Exception as e:
log.error(f'Error getting OAuth session by provider and user ID: {e}')
return None
def get_sessions_by_user_id(self, user_id: str, db: Optional[Session] = None) -> List[OAuthSessionModel]:
async def get_sessions_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> List[OAuthSessionModel]:
"""Get all OAuth sessions for a user"""
try:
with get_db_context(db) as db:
sessions = db.query(OAuthSession).filter_by(user_id=user_id).all()
async with get_async_db_context(db) as db:
result = await db.execute(select(OAuthSession).filter_by(user_id=user_id))
sessions = result.scalars().all()
results = []
for session in sessions:
try:
db.expunge(session)
session.token = self._decrypt_token(session.token)
results.append(OAuthSessionModel.model_validate(session))
results.append(OAuthSessionModel(
id=session.id,
user_id=session.user_id,
provider=session.provider,
token=self._decrypt_token(session.token),
expires_at=session.expires_at,
created_at=session.created_at,
updated_at=session.updated_at,
))
except Exception as e:
log.warning(
f'Skipping OAuth session {session.id} due to decryption failure, deleting corrupted session: {type(e).__name__}: {e}'
)
db.query(OAuthSession).filter_by(id=session.id).delete()
db.commit()
await db.execute(delete(OAuthSession).filter_by(id=session.id))
await db.commit()
return results
@ -221,62 +257,69 @@ class OAuthSessionTable:
log.error(f'Error getting OAuth sessions by user ID: {e}')
return []
def update_session_by_id(
self, session_id: str, token: dict, db: Optional[Session] = None
async def update_session_by_id(
self, session_id: str, token: dict, db: Optional[AsyncSession] = None
) -> Optional[OAuthSessionModel]:
"""Update OAuth session tokens"""
try:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
current_time = int(time.time())
db.query(OAuthSession).filter_by(id=session_id).update(
{
'token': self._encrypt_token(token),
'expires_at': token.get('expires_at'),
'updated_at': current_time,
}
await db.execute(
update(OAuthSession).filter_by(id=session_id).values(
token=self._encrypt_token(token),
expires_at=token.get('expires_at'),
updated_at=current_time,
)
)
db.commit()
session = db.query(OAuthSession).filter_by(id=session_id).first()
await db.commit()
result = await db.execute(select(OAuthSession).filter_by(id=session_id))
session = result.scalars().first()
if session:
db.expunge(session)
session.token = self._decrypt_token(session.token)
return OAuthSessionModel.model_validate(session)
return OAuthSessionModel(
id=session.id,
user_id=session.user_id,
provider=session.provider,
token=self._decrypt_token(session.token),
expires_at=session.expires_at,
created_at=session.created_at,
updated_at=session.updated_at,
)
return None
except Exception as e:
log.error(f'Error updating OAuth session tokens: {e}')
return None
def delete_session_by_id(self, session_id: str, db: Optional[Session] = None) -> bool:
async def delete_session_by_id(self, session_id: str, db: Optional[AsyncSession] = None) -> bool:
"""Delete an OAuth session"""
try:
with get_db_context(db) as db:
result = db.query(OAuthSession).filter_by(id=session_id).delete()
db.commit()
return result > 0
async with get_async_db_context(db) as db:
result = await db.execute(delete(OAuthSession).filter_by(id=session_id))
await db.commit()
return result.rowcount > 0
except Exception as e:
log.error(f'Error deleting OAuth session: {e}')
return False
def delete_sessions_by_user_id(self, user_id: str, db: Optional[Session] = None) -> bool:
async def delete_sessions_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> bool:
"""Delete all OAuth sessions for a user"""
try:
with get_db_context(db) as db:
result = db.query(OAuthSession).filter_by(user_id=user_id).delete()
db.commit()
async with get_async_db_context(db) as db:
await db.execute(delete(OAuthSession).filter_by(user_id=user_id))
await db.commit()
return True
except Exception as e:
log.error(f'Error deleting OAuth sessions by user ID: {e}')
return False
def delete_sessions_by_provider(self, provider: str, db: Optional[Session] = None) -> bool:
async def delete_sessions_by_provider(self, provider: str, db: Optional[AsyncSession] = None) -> bool:
"""Delete all OAuth sessions for a provider"""
try:
with get_db_context(db) as db:
db.query(OAuthSession).filter_by(provider=provider).delete()
db.commit()
async with get_async_db_context(db) as db:
await db.execute(delete(OAuthSession).filter_by(provider=provider))
await db.commit()
return True
except Exception as e:
log.error(f'Error deleting OAuth sessions by provider {provider}: {e}')

View file

@ -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

View file

@ -2,16 +2,17 @@ import time
import uuid
from typing import Optional
from sqlalchemy.orm import Session
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
from sqlalchemy import select, delete, update, or_, func, cast, String
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import Base, JSONField, get_async_db_context
from open_webui.models.groups import Groups
from open_webui.models.users import Users, UserResponse
from open_webui.models.users import Users, User, UserModel, UserResponse
from open_webui.models.prompt_history import PromptHistories
from open_webui.models.access_grants import AccessGrantModel, AccessGrants
from pydantic import BaseModel, ConfigDict, Field
from sqlalchemy import BigInteger, Boolean, Column, String, Text, JSON, or_, func, cast
from sqlalchemy import BigInteger, Boolean, Column, Text, JSON
####################
# Prompts DB Schema
@ -92,23 +93,23 @@ class PromptForm(BaseModel):
class PromptsTable:
def _get_access_grants(self, prompt_id: str, db: Optional[Session] = None) -> list[AccessGrantModel]:
return AccessGrants.get_grants_by_resource('prompt', prompt_id, db=db)
async def _get_access_grants(self, prompt_id: str, db: Optional[AsyncSession] = None) -> list[AccessGrantModel]:
return await AccessGrants.get_grants_by_resource('prompt', prompt_id, db=db)
def _to_prompt_model(
async def _to_prompt_model(
self,
prompt: Prompt,
access_grants: Optional[list[AccessGrantModel]] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> PromptModel:
prompt_data = PromptModel.model_validate(prompt).model_dump(exclude={'access_grants'})
prompt_data['access_grants'] = (
access_grants if access_grants is not None else self._get_access_grants(prompt_data['id'], db=db)
access_grants if access_grants is not None else await self._get_access_grants(prompt_data['id'], db=db)
)
return PromptModel.model_validate(prompt_data)
def insert_new_prompt(
self, user_id: str, form_data: PromptForm, db: Optional[Session] = None
async def insert_new_prompt(
self, user_id: str, form_data: PromptForm, db: Optional[AsyncSession] = None
) -> Optional[PromptModel]:
now = int(time.time())
prompt_id = str(uuid.uuid4())
@ -129,15 +130,15 @@ class PromptsTable:
)
try:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
result = Prompt(**prompt.model_dump(exclude={'access_grants'}))
db.add(result)
db.commit()
db.refresh(result)
AccessGrants.set_access_grants('prompt', prompt_id, form_data.access_grants, db=db)
await db.commit()
await db.refresh(result)
await AccessGrants.set_access_grants('prompt', prompt_id, form_data.access_grants, db=db)
if result:
current_access_grants = self._get_access_grants(prompt_id, db=db)
current_access_grants = await self._get_access_grants(prompt_id, db=db)
snapshot = {
'name': form_data.name,
'content': form_data.content,
@ -148,7 +149,7 @@ class PromptsTable:
'access_grants': [grant.model_dump() for grant in current_access_grants],
}
history_entry = PromptHistories.create_history_entry(
history_entry = await PromptHistories.create_history_entry(
prompt_id=prompt_id,
snapshot=snapshot,
user_id=user_id,
@ -160,46 +161,51 @@ class PromptsTable:
# Set the initial version as the production version
if history_entry:
result.version_id = history_entry.id
db.commit()
db.refresh(result)
await db.commit()
await db.refresh(result)
return self._to_prompt_model(result, db=db)
return await self._to_prompt_model(result, db=db)
else:
return None
except Exception:
return None
def get_prompt_by_id(self, prompt_id: str, db: Optional[Session] = None) -> Optional[PromptModel]:
async def get_prompt_by_id(self, prompt_id: str, db: Optional[AsyncSession] = None) -> Optional[PromptModel]:
"""Get prompt by UUID."""
try:
with get_db_context(db) as db:
prompt = db.query(Prompt).filter_by(id=prompt_id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(Prompt).filter_by(id=prompt_id))
prompt = result.scalars().first()
if prompt:
return self._to_prompt_model(prompt, db=db)
return await self._to_prompt_model(prompt, db=db)
return None
except Exception:
return None
def get_prompt_by_command(self, command: str, db: Optional[Session] = None) -> Optional[PromptModel]:
async def get_prompt_by_command(self, command: str, db: Optional[AsyncSession] = None) -> Optional[PromptModel]:
try:
with get_db_context(db) as db:
prompt = db.query(Prompt).filter_by(command=command).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(Prompt).filter_by(command=command))
prompt = result.scalars().first()
if prompt:
return self._to_prompt_model(prompt, db=db)
return await self._to_prompt_model(prompt, db=db)
return None
except Exception:
return None
def get_prompts(self, db: Optional[Session] = None) -> list[PromptUserResponse]:
with get_db_context(db) as db:
all_prompts = db.query(Prompt).filter(Prompt.is_active == True).order_by(Prompt.updated_at.desc()).all()
async def get_prompts(self, db: Optional[AsyncSession] = None) -> list[PromptUserResponse]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(Prompt).filter(Prompt.is_active == True).order_by(Prompt.updated_at.desc())
)
all_prompts = result.scalars().all()
user_ids = list(set(prompt.user_id for prompt in all_prompts))
prompt_ids = [prompt.id for prompt in all_prompts]
users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
users = await Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
users_dict = {user.id: user for user in users}
grants_map = AccessGrants.get_grants_by_resources('prompt', prompt_ids, db=db)
grants_map = await AccessGrants.get_grants_by_resources('prompt', prompt_ids, db=db)
prompts = []
for prompt in all_prompts:
@ -207,11 +213,11 @@ class PromptsTable:
prompts.append(
PromptUserResponse.model_validate(
{
**self._to_prompt_model(
**(await self._to_prompt_model(
prompt,
access_grants=grants_map.get(prompt.id, []),
db=db,
).model_dump(),
)).model_dump(),
'user': user.model_dump() if user else None,
}
)
@ -219,44 +225,44 @@ class PromptsTable:
return prompts
def get_prompts_by_user_id(
self, user_id: str, permission: str = 'write', db: Optional[Session] = None
async def get_prompts_by_user_id(
self, user_id: str, permission: str = 'write', db: Optional[AsyncSession] = None
) -> list[PromptUserResponse]:
prompts = self.get_prompts(db=db)
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id, db=db)}
prompts = await self.get_prompts(db=db)
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
user_group_ids = {group.id for group in user_groups}
return [
prompt
for prompt in prompts
if prompt.user_id == user_id
or AccessGrants.has_access(
result = []
for prompt in prompts:
if prompt.user_id == user_id:
result.append(prompt)
elif await AccessGrants.has_access(
user_id=user_id,
resource_type='prompt',
resource_id=prompt.id,
permission=permission,
user_group_ids=user_group_ids,
db=db,
)
]
):
result.append(prompt)
return result
def search_prompts(
async def search_prompts(
self,
user_id: str,
filter: dict = {},
skip: int = 0,
limit: int = 30,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> PromptListResponse:
with get_db_context(db) as db:
from open_webui.models.users import User, UserModel
async with get_async_db_context(db) as db:
# Join with User table for user filtering and sorting
query = db.query(Prompt, User).outerjoin(User, User.id == Prompt.user_id)
stmt = select(Prompt, User).outerjoin(User, User.id == Prompt.user_id)
if filter:
query_key = filter.get('query')
if query_key:
query = query.filter(
stmt = stmt.filter(
or_(
Prompt.name.ilike(f'%{query_key}%'),
Prompt.command.ilike(f'%{query_key}%'),
@ -268,14 +274,14 @@ class PromptsTable:
view_option = filter.get('view_option')
if view_option == 'created':
query = query.filter(Prompt.user_id == user_id)
stmt = stmt.filter(Prompt.user_id == user_id)
elif view_option == 'shared':
query = query.filter(Prompt.user_id != user_id)
stmt = stmt.filter(Prompt.user_id != user_id)
# Apply access grant filtering
query = AccessGrants.has_permission_filter(
stmt = AccessGrants.has_permission_filter(
db=db,
query=query,
query=stmt,
DocumentModel=Prompt,
filter=filter,
resource_type='prompt',
@ -287,75 +293,80 @@ class PromptsTable:
# Search for tag in JSON array field
like_pattern = f'%"{tag.lower()}"%'
tags_text = func.lower(cast(Prompt.tags, String))
query = query.filter(tags_text.like(like_pattern))
stmt = stmt.filter(tags_text.like(like_pattern))
order_by = filter.get('order_by')
direction = filter.get('direction')
if order_by == 'name':
if direction == 'asc':
query = query.order_by(Prompt.name.asc())
stmt = stmt.order_by(Prompt.name.asc())
else:
query = query.order_by(Prompt.name.desc())
stmt = stmt.order_by(Prompt.name.desc())
elif order_by == 'created_at':
if direction == 'asc':
query = query.order_by(Prompt.created_at.asc())
stmt = stmt.order_by(Prompt.created_at.asc())
else:
query = query.order_by(Prompt.created_at.desc())
stmt = stmt.order_by(Prompt.created_at.desc())
elif order_by == 'updated_at':
if direction == 'asc':
query = query.order_by(Prompt.updated_at.asc())
stmt = stmt.order_by(Prompt.updated_at.asc())
else:
query = query.order_by(Prompt.updated_at.desc())
stmt = stmt.order_by(Prompt.updated_at.desc())
else:
query = query.order_by(Prompt.updated_at.desc())
stmt = stmt.order_by(Prompt.updated_at.desc())
else:
query = query.order_by(Prompt.updated_at.desc())
stmt = stmt.order_by(Prompt.updated_at.desc())
# Count BEFORE pagination
total = query.count()
count_result = await db.execute(
select(func.count()).select_from(stmt.subquery())
)
total = count_result.scalar()
if skip:
query = query.offset(skip)
stmt = stmt.offset(skip)
if limit:
query = query.limit(limit)
stmt = stmt.limit(limit)
items = query.all()
result = await db.execute(stmt)
items = result.all()
prompt_ids = [prompt.id for prompt, _ in items]
grants_map = AccessGrants.get_grants_by_resources('prompt', prompt_ids, db=db)
grants_map = await AccessGrants.get_grants_by_resources('prompt', prompt_ids, db=db)
prompts = []
for prompt, user in items:
prompts.append(
PromptUserResponse(
**self._to_prompt_model(
**(await self._to_prompt_model(
prompt,
access_grants=grants_map.get(prompt.id, []),
db=db,
).model_dump(),
)).model_dump(),
user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
)
)
return PromptListResponse(items=prompts, total=total)
def update_prompt_by_command(
async def update_prompt_by_command(
self,
command: str,
form_data: PromptForm,
user_id: str,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[PromptModel]:
try:
with get_db_context(db) as db:
prompt = db.query(Prompt).filter_by(command=command).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(Prompt).filter_by(command=command))
prompt = result.scalars().first()
if not prompt:
return None
latest_history = PromptHistories.get_latest_history_entry(prompt.id, db=db)
latest_history = await PromptHistories.get_latest_history_entry(prompt.id, db=db)
parent_id = latest_history.id if latest_history else None
current_access_grants = self._get_access_grants(prompt.id, db=db)
current_access_grants = await self._get_access_grants(prompt.id, db=db)
# Check if content changed to decide on history creation
content_changed = (
@ -371,10 +382,10 @@ class PromptsTable:
prompt.meta = form_data.meta or prompt.meta
prompt.updated_at = int(time.time())
if form_data.access_grants is not None:
AccessGrants.set_access_grants('prompt', prompt.id, form_data.access_grants, db=db)
current_access_grants = self._get_access_grants(prompt.id, db=db)
await AccessGrants.set_access_grants('prompt', prompt.id, form_data.access_grants, db=db)
current_access_grants = await self._get_access_grants(prompt.id, db=db)
db.commit()
await db.commit()
# Create history entry only if content changed
if content_changed:
@ -387,7 +398,7 @@ class PromptsTable:
'access_grants': [grant.model_dump() for grant in current_access_grants],
}
history_entry = PromptHistories.create_history_entry(
history_entry = await PromptHistories.create_history_entry(
prompt_id=prompt.id,
snapshot=snapshot,
user_id=user_id,
@ -399,28 +410,29 @@ class PromptsTable:
# Set as production if flag is True (default)
if form_data.is_production and history_entry:
prompt.version_id = history_entry.id
db.commit()
await db.commit()
return self._to_prompt_model(prompt, db=db)
return await self._to_prompt_model(prompt, db=db)
except Exception:
return None
def update_prompt_by_id(
async def update_prompt_by_id(
self,
prompt_id: str,
form_data: PromptForm,
user_id: str,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[PromptModel]:
try:
with get_db_context(db) as db:
prompt = db.query(Prompt).filter_by(id=prompt_id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(Prompt).filter_by(id=prompt_id))
prompt = result.scalars().first()
if not prompt:
return None
latest_history = PromptHistories.get_latest_history_entry(prompt.id, db=db)
latest_history = await PromptHistories.get_latest_history_entry(prompt.id, db=db)
parent_id = latest_history.id if latest_history else None
current_access_grants = self._get_access_grants(prompt.id, db=db)
current_access_grants = await self._get_access_grants(prompt.id, db=db)
# Check if content changed to decide on history creation
content_changed = (
@ -442,12 +454,12 @@ class PromptsTable:
prompt.tags = form_data.tags
if form_data.access_grants is not None:
AccessGrants.set_access_grants('prompt', prompt.id, form_data.access_grants, db=db)
current_access_grants = self._get_access_grants(prompt.id, db=db)
await AccessGrants.set_access_grants('prompt', prompt.id, form_data.access_grants, db=db)
current_access_grants = await self._get_access_grants(prompt.id, db=db)
prompt.updated_at = int(time.time())
db.commit()
await db.commit()
# Create history entry only if content changed
if content_changed:
@ -461,7 +473,7 @@ class PromptsTable:
'access_grants': [grant.model_dump() for grant in current_access_grants],
}
history_entry = PromptHistories.create_history_entry(
history_entry = await PromptHistories.create_history_entry(
prompt_id=prompt.id,
snapshot=snapshot,
user_id=user_id,
@ -473,24 +485,25 @@ class PromptsTable:
# Set as production if flag is True (default)
if form_data.is_production and history_entry:
prompt.version_id = history_entry.id
db.commit()
await db.commit()
return self._to_prompt_model(prompt, db=db)
return await self._to_prompt_model(prompt, db=db)
except Exception:
return None
def update_prompt_metadata(
async def update_prompt_metadata(
self,
prompt_id: str,
name: str,
command: str,
tags: Optional[list[str]] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[PromptModel]:
"""Update only name, command, and tags (no history created)."""
try:
with get_db_context(db) as db:
prompt = db.query(Prompt).filter_by(id=prompt_id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(Prompt).filter_by(id=prompt_id))
prompt = result.scalars().first()
if not prompt:
return None
@ -501,26 +514,27 @@ class PromptsTable:
prompt.tags = tags
prompt.updated_at = int(time.time())
db.commit()
await db.commit()
return self._to_prompt_model(prompt, db=db)
return await self._to_prompt_model(prompt, db=db)
except Exception:
return None
def update_prompt_version(
async def update_prompt_version(
self,
prompt_id: str,
version_id: str,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[PromptModel]:
"""Set the active version of a prompt and restore content from that version's snapshot."""
try:
with get_db_context(db) as db:
prompt = db.query(Prompt).filter_by(id=prompt_id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(Prompt).filter_by(id=prompt_id))
prompt = result.scalars().first()
if not prompt:
return None
history_entry = PromptHistories.get_history_entry_by_id(version_id, db=db)
history_entry = await PromptHistories.get_history_entry_by_id(version_id, db=db)
if not history_entry:
return None
@ -537,63 +551,67 @@ class PromptsTable:
prompt.version_id = version_id
prompt.updated_at = int(time.time())
db.commit()
await db.commit()
return self._to_prompt_model(prompt, db=db)
return await self._to_prompt_model(prompt, db=db)
except Exception:
return None
def toggle_prompt_active(self, prompt_id: str, db: Optional[Session] = None) -> Optional[PromptModel]:
async def toggle_prompt_active(self, prompt_id: str, db: Optional[AsyncSession] = None) -> Optional[PromptModel]:
"""Toggle the is_active flag on a prompt."""
try:
with get_db_context(db) as db:
prompt = db.query(Prompt).filter_by(id=prompt_id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(Prompt).filter_by(id=prompt_id))
prompt = result.scalars().first()
if prompt:
prompt.is_active = not prompt.is_active
prompt.updated_at = int(time.time())
db.commit()
db.refresh(prompt)
return self._to_prompt_model(prompt, db=db)
await db.commit()
await db.refresh(prompt)
return await self._to_prompt_model(prompt, db=db)
return None
except Exception:
return None
def delete_prompt_by_command(self, command: str, db: Optional[Session] = None) -> bool:
async def delete_prompt_by_command(self, command: str, db: Optional[AsyncSession] = None) -> bool:
"""Permanently delete a prompt and its history."""
try:
with get_db_context(db) as db:
prompt = db.query(Prompt).filter_by(command=command).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(Prompt).filter_by(command=command))
prompt = result.scalars().first()
if prompt:
PromptHistories.delete_history_by_prompt_id(prompt.id, db=db)
AccessGrants.revoke_all_access('prompt', prompt.id, db=db)
await PromptHistories.delete_history_by_prompt_id(prompt.id, db=db)
await AccessGrants.revoke_all_access('prompt', prompt.id, db=db)
db.delete(prompt)
db.commit()
await db.delete(prompt)
await db.commit()
return True
return False
except Exception:
return False
def delete_prompt_by_id(self, prompt_id: str, db: Optional[Session] = None) -> bool:
async def delete_prompt_by_id(self, prompt_id: str, db: Optional[AsyncSession] = None) -> bool:
"""Permanently delete a prompt and its history."""
try:
with get_db_context(db) as db:
prompt = db.query(Prompt).filter_by(id=prompt_id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(Prompt).filter_by(id=prompt_id))
prompt = result.scalars().first()
if prompt:
PromptHistories.delete_history_by_prompt_id(prompt.id, db=db)
AccessGrants.revoke_all_access('prompt', prompt.id, db=db)
await PromptHistories.delete_history_by_prompt_id(prompt.id, db=db)
await AccessGrants.revoke_all_access('prompt', prompt.id, db=db)
db.delete(prompt)
db.commit()
await db.delete(prompt)
await db.commit()
return True
return False
except Exception:
return False
def get_tags(self, db: Optional[Session] = None) -> list[str]:
async def get_tags(self, db: Optional[AsyncSession] = None) -> list[str]:
try:
with get_db_context(db) as db:
prompts = db.query(Prompt).filter_by(is_active=True).all()
async with get_async_db_context(db) as db:
result = await db.execute(select(Prompt).filter_by(is_active=True))
prompts = result.scalars().all()
tags = set()
for prompt in prompts:
if prompt.tags:

View file

@ -2,14 +2,15 @@ import logging
import time
from typing import Optional
from sqlalchemy.orm import Session
from open_webui.internal.db import Base, get_db, get_db_context
from open_webui.models.users import Users, UserResponse
from sqlalchemy import select, delete, update, or_
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import Base, get_async_db_context
from open_webui.models.users import Users, User, UserModel, UserResponse
from open_webui.models.groups import Groups
from open_webui.models.access_grants import AccessGrantModel, AccessGrants
from pydantic import BaseModel, ConfigDict, Field
from sqlalchemy import JSON, BigInteger, Boolean, Column, String, Text, or_
from sqlalchemy import JSON, BigInteger, Boolean, Column, String, Text, func
log = logging.getLogger(__name__)
@ -105,28 +106,28 @@ class SkillAccessListResponse(BaseModel):
class SkillsTable:
def _get_access_grants(self, skill_id: str, db: Optional[Session] = None) -> list[AccessGrantModel]:
return AccessGrants.get_grants_by_resource('skill', skill_id, db=db)
async def _get_access_grants(self, skill_id: str, db: Optional[AsyncSession] = None) -> list[AccessGrantModel]:
return await AccessGrants.get_grants_by_resource('skill', skill_id, db=db)
def _to_skill_model(
async def _to_skill_model(
self,
skill: Skill,
access_grants: Optional[list[AccessGrantModel]] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> SkillModel:
skill_data = SkillModel.model_validate(skill).model_dump(exclude={'access_grants'})
skill_data['access_grants'] = (
access_grants if access_grants is not None else self._get_access_grants(skill_data['id'], db=db)
access_grants if access_grants is not None else await self._get_access_grants(skill_data['id'], db=db)
)
return SkillModel.model_validate(skill_data)
def insert_new_skill(
async def insert_new_skill(
self,
user_id: str,
form_data: SkillForm,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[SkillModel]:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
try:
result = Skill(
**{
@ -137,43 +138,45 @@ class SkillsTable:
}
)
db.add(result)
db.commit()
db.refresh(result)
AccessGrants.set_access_grants('skill', result.id, form_data.access_grants, db=db)
await db.commit()
await db.refresh(result)
await AccessGrants.set_access_grants('skill', result.id, form_data.access_grants, db=db)
if result:
return self._to_skill_model(result, db=db)
return await self._to_skill_model(result, db=db)
else:
return None
except Exception as e:
log.exception(f'Error creating a new skill: {e}')
return None
def get_skill_by_id(self, id: str, db: Optional[Session] = None) -> Optional[SkillModel]:
async def get_skill_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[SkillModel]:
try:
with get_db_context(db) as db:
skill = db.get(Skill, id)
return self._to_skill_model(skill, db=db) if skill else None
async with get_async_db_context(db) as db:
skill = await db.get(Skill, id)
return await self._to_skill_model(skill, db=db) if skill else None
except Exception:
return None
def get_skill_by_name(self, name: str, db: Optional[Session] = None) -> Optional[SkillModel]:
async def get_skill_by_name(self, name: str, db: Optional[AsyncSession] = None) -> Optional[SkillModel]:
try:
with get_db_context(db) as db:
skill = db.query(Skill).filter_by(name=name).first()
return self._to_skill_model(skill, db=db) if skill else None
async with get_async_db_context(db) as db:
result = await db.execute(select(Skill).filter_by(name=name))
skill = result.scalars().first()
return await self._to_skill_model(skill, db=db) if skill else None
except Exception:
return None
def get_skills(self, db: Optional[Session] = None) -> list[SkillUserModel]:
with get_db_context(db) as db:
all_skills = db.query(Skill).order_by(Skill.updated_at.desc()).all()
async def get_skills(self, db: Optional[AsyncSession] = None) -> list[SkillUserModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(Skill).order_by(Skill.updated_at.desc()))
all_skills = result.scalars().all()
user_ids = list(set(skill.user_id for skill in all_skills))
skill_ids = [skill.id for skill in all_skills]
users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
users = await Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
users_dict = {user.id: user for user in users}
grants_map = AccessGrants.get_grants_by_resources('skill', skill_ids, db=db)
grants_map = await AccessGrants.get_grants_by_resources('skill', skill_ids, db=db)
skills = []
for skill in all_skills:
@ -181,56 +184,56 @@ class SkillsTable:
skills.append(
SkillUserModel.model_validate(
{
**self._to_skill_model(
**(await self._to_skill_model(
skill,
access_grants=grants_map.get(skill.id, []),
db=db,
).model_dump(),
)).model_dump(),
'user': user.model_dump() if user else None,
}
)
)
return skills
def get_skills_by_user_id(
self, user_id: str, permission: str = 'write', db: Optional[Session] = None
async def get_skills_by_user_id(
self, user_id: str, permission: str = 'write', db: Optional[AsyncSession] = None
) -> list[SkillUserModel]:
skills = self.get_skills(db=db)
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id, db=db)}
skills = await self.get_skills(db=db)
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
user_group_ids = {group.id for group in user_groups}
return [
skill
for skill in skills
if skill.user_id == user_id
or AccessGrants.has_access(
result = []
for skill in skills:
if skill.user_id == user_id:
result.append(skill)
elif await AccessGrants.has_access(
user_id=user_id,
resource_type='skill',
resource_id=skill.id,
permission=permission,
user_group_ids=user_group_ids,
db=db,
)
]
):
result.append(skill)
return result
def search_skills(
async def search_skills(
self,
user_id: str,
filter: dict = {},
skip: int = 0,
limit: int = 30,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> SkillListResponse:
try:
with get_db_context(db) as db:
from open_webui.models.users import User, UserModel
async with get_async_db_context(db) as db:
# Join with User table for user filtering
query = db.query(Skill, User).outerjoin(User, User.id == Skill.user_id)
stmt = select(Skill, User).outerjoin(User, User.id == Skill.user_id)
if filter:
query_key = filter.get('query')
if query_key:
query = query.filter(
stmt = stmt.filter(
or_(
Skill.name.ilike(f'%{query_key}%'),
Skill.description.ilike(f'%{query_key}%'),
@ -242,44 +245,48 @@ class SkillsTable:
view_option = filter.get('view_option')
if view_option == 'created':
query = query.filter(Skill.user_id == user_id)
stmt = stmt.filter(Skill.user_id == user_id)
elif view_option == 'shared':
query = query.filter(Skill.user_id != user_id)
stmt = stmt.filter(Skill.user_id != user_id)
# Apply access grant filtering
query = AccessGrants.has_permission_filter(
stmt = AccessGrants.has_permission_filter(
db=db,
query=query,
query=stmt,
DocumentModel=Skill,
filter=filter,
resource_type='skill',
permission='read',
)
query = query.order_by(Skill.updated_at.desc())
stmt = stmt.order_by(Skill.updated_at.desc())
# Count BEFORE pagination
total = query.count()
count_result = await db.execute(
select(func.count()).select_from(stmt.subquery())
)
total = count_result.scalar()
if skip:
query = query.offset(skip)
stmt = stmt.offset(skip)
if limit:
query = query.limit(limit)
stmt = stmt.limit(limit)
items = query.all()
result = await db.execute(stmt)
items = result.all()
skill_ids = [skill.id for skill, _ in items]
grants_map = AccessGrants.get_grants_by_resources('skill', skill_ids, db=db)
grants_map = await AccessGrants.get_grants_by_resources('skill', skill_ids, db=db)
skills = []
for skill, user in items:
skills.append(
SkillUserResponse(
**self._to_skill_model(
**(await self._to_skill_model(
skill,
access_grants=grants_map.get(skill.id, []),
db=db,
).model_dump(),
)).model_dump(),
user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
)
)
@ -289,43 +296,44 @@ class SkillsTable:
log.exception(f'Error searching skills: {e}')
return SkillListResponse(items=[], total=0)
def update_skill_by_id(self, id: str, updated: dict, db: Optional[Session] = None) -> Optional[SkillModel]:
async def update_skill_by_id(self, id: str, updated: dict, db: Optional[AsyncSession] = None) -> Optional[SkillModel]:
try:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
access_grants = updated.pop('access_grants', None)
db.query(Skill).filter_by(id=id).update({**updated, 'updated_at': int(time.time())})
db.commit()
await db.execute(update(Skill).filter_by(id=id).values(**updated, updated_at=int(time.time())))
await db.commit()
if access_grants is not None:
AccessGrants.set_access_grants('skill', id, access_grants, db=db)
await AccessGrants.set_access_grants('skill', id, access_grants, db=db)
skill = db.query(Skill).get(id)
db.refresh(skill)
return self._to_skill_model(skill, db=db)
skill = await db.get(Skill, id)
await db.refresh(skill)
return await self._to_skill_model(skill, db=db)
except Exception:
return None
def toggle_skill_by_id(self, id: str, db: Optional[Session] = None) -> Optional[SkillModel]:
with get_db_context(db) as db:
async def toggle_skill_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[SkillModel]:
async with get_async_db_context(db) as db:
try:
skill = db.query(Skill).filter_by(id=id).first()
result = await db.execute(select(Skill).filter_by(id=id))
skill = result.scalars().first()
if not skill:
return None
skill.is_active = not skill.is_active
skill.updated_at = int(time.time())
db.commit()
db.refresh(skill)
await db.commit()
await db.refresh(skill)
return self._to_skill_model(skill, db=db)
return await self._to_skill_model(skill, db=db)
except Exception:
return None
def delete_skill_by_id(self, id: str, db: Optional[Session] = None) -> bool:
async def delete_skill_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
try:
with get_db_context(db) as db:
AccessGrants.revoke_all_access('skill', id, db=db)
db.query(Skill).filter_by(id=id).delete()
db.commit()
async with get_async_db_context(db) as db:
await AccessGrants.revoke_all_access('skill', id, db=db)
await db.execute(delete(Skill).filter_by(id=id))
await db.commit()
return True
except Exception:

View file

@ -3,8 +3,9 @@ import time
import uuid
from typing import Optional
from sqlalchemy.orm import Session
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
from sqlalchemy import select, delete
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import Base, JSONField, get_async_db_context
from pydantic import BaseModel, ConfigDict
@ -53,15 +54,15 @@ class TagChatIdForm(BaseModel):
class TagTable:
def insert_new_tag(self, name: str, user_id: str, db: Optional[Session] = None) -> Optional[TagModel]:
with get_db_context(db) as db:
async def insert_new_tag(self, name: str, user_id: str, db: Optional[AsyncSession] = None) -> Optional[TagModel]:
async with get_async_db_context(db) as db:
id = name.replace(' ', '_').lower()
tag = TagModel(**{'id': id, 'user_id': user_id, 'name': name})
try:
result = Tag(**tag.model_dump())
db.add(result)
db.commit()
db.refresh(result)
await db.commit()
await db.refresh(result)
if result:
return TagModel.model_validate(result)
else:
@ -70,64 +71,65 @@ class TagTable:
log.exception(f'Error inserting a new tag: {e}')
return None
def get_tag_by_name_and_user_id(self, name: str, user_id: str, db: Optional[Session] = None) -> Optional[TagModel]:
async def get_tag_by_name_and_user_id(self, name: str, user_id: str, db: Optional[AsyncSession] = None) -> Optional[TagModel]:
try:
id = name.replace(' ', '_').lower()
with get_db_context(db) as db:
tag = db.query(Tag).filter_by(id=id, user_id=user_id).first()
return TagModel.model_validate(tag)
async with get_async_db_context(db) as db:
result = await db.execute(select(Tag).filter_by(id=id, user_id=user_id))
tag = result.scalars().first()
return TagModel.model_validate(tag) if tag else None
except Exception:
return None
def get_tags_by_user_id(self, user_id: str, db: Optional[Session] = None) -> list[TagModel]:
with get_db_context(db) as db:
return [TagModel.model_validate(tag) for tag in (db.query(Tag).filter_by(user_id=user_id).all())]
async def get_tags_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[TagModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(Tag).filter_by(user_id=user_id))
return [TagModel.model_validate(tag) for tag in result.scalars().all()]
def get_tags_by_ids_and_user_id(self, ids: list[str], user_id: str, db: Optional[Session] = None) -> list[TagModel]:
with get_db_context(db) as db:
return [
TagModel.model_validate(tag)
for tag in (db.query(Tag).filter(Tag.id.in_(ids), Tag.user_id == user_id).all())
]
async def get_tags_by_ids_and_user_id(self, ids: list[str], user_id: str, db: Optional[AsyncSession] = None) -> list[TagModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(Tag).filter(Tag.id.in_(ids), Tag.user_id == user_id))
return [TagModel.model_validate(tag) for tag in result.scalars().all()]
def delete_tag_by_name_and_user_id(self, name: str, user_id: str, db: Optional[Session] = None) -> bool:
async def delete_tag_by_name_and_user_id(self, name: str, user_id: str, db: Optional[AsyncSession] = None) -> bool:
try:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
id = name.replace(' ', '_').lower()
res = db.query(Tag).filter_by(id=id, user_id=user_id).delete()
log.debug(f'res: {res}')
db.commit()
result = await db.execute(delete(Tag).filter_by(id=id, user_id=user_id))
log.debug(f'res: {result.rowcount}')
await db.commit()
return True
except Exception as e:
log.error(f'delete_tag: {e}')
return False
def delete_tags_by_ids_and_user_id(self, ids: list[str], user_id: str, db: Optional[Session] = None) -> bool:
async def delete_tags_by_ids_and_user_id(self, ids: list[str], user_id: str, db: Optional[AsyncSession] = None) -> bool:
"""Delete all tags whose id is in *ids* for the given user, in one query."""
if not ids:
return True
try:
with get_db_context(db) as db:
db.query(Tag).filter(Tag.id.in_(ids), Tag.user_id == user_id).delete(synchronize_session=False)
db.commit()
async with get_async_db_context(db) as db:
await db.execute(delete(Tag).filter(Tag.id.in_(ids), Tag.user_id == user_id))
await db.commit()
return True
except Exception as e:
log.error(f'delete_tags_by_ids: {e}')
return False
def ensure_tags_exist(self, names: list[str], user_id: str, db: Optional[Session] = None) -> None:
async def ensure_tags_exist(self, names: list[str], user_id: str, db: Optional[AsyncSession] = None) -> None:
"""Create tag rows for any *names* that don't already exist for *user_id*."""
if not names:
return
ids = [n.replace(' ', '_').lower() for n in names]
with get_db_context(db) as db:
existing = {t.id for t in db.query(Tag.id).filter(Tag.id.in_(ids), Tag.user_id == user_id).all()}
async with get_async_db_context(db) as db:
result = await db.execute(select(Tag.id).filter(Tag.id.in_(ids), Tag.user_id == user_id))
existing = {row[0] for row in result.all()}
new_tags = [
Tag(id=tag_id, name=name, user_id=user_id) for tag_id, name in zip(ids, names) if tag_id not in existing
]
if new_tags:
db.add_all(new_tags)
db.commit()
await db.commit()
Tags = TagTable()

View file

@ -2,8 +2,9 @@ import logging
import time
from typing import Optional
from sqlalchemy.orm import Session, defer
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
from sqlalchemy import select, delete, update
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import Base, JSONField, get_async_db_context
from open_webui.models.users import Users, UserResponse
from open_webui.models.groups import Groups
from open_webui.models.access_grants import AccessGrantModel, AccessGrants
@ -97,29 +98,29 @@ class ToolValves(BaseModel):
class ToolsTable:
def _get_access_grants(self, tool_id: str, db: Optional[Session] = None) -> list[AccessGrantModel]:
return AccessGrants.get_grants_by_resource('tool', tool_id, db=db)
async def _get_access_grants(self, tool_id: str, db: Optional[AsyncSession] = None) -> list[AccessGrantModel]:
return await AccessGrants.get_grants_by_resource('tool', tool_id, db=db)
def _to_tool_model(
async def _to_tool_model(
self,
tool: Tool,
access_grants: Optional[list[AccessGrantModel]] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> ToolModel:
tool_data = ToolModel.model_validate(tool).model_dump(exclude={'access_grants'})
tool_data['access_grants'] = (
access_grants if access_grants is not None else self._get_access_grants(tool_data['id'], db=db)
access_grants if access_grants is not None else await self._get_access_grants(tool_data['id'], db=db)
)
return ToolModel.model_validate(tool_data)
def insert_new_tool(
async def insert_new_tool(
self,
user_id: str,
form_data: ToolForm,
specs: list[dict],
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[ToolModel]:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
try:
result = Tool(
**{
@ -131,38 +132,39 @@ class ToolsTable:
}
)
db.add(result)
db.commit()
db.refresh(result)
AccessGrants.set_access_grants('tool', result.id, form_data.access_grants, db=db)
await db.commit()
await db.refresh(result)
await AccessGrants.set_access_grants('tool', result.id, form_data.access_grants, db=db)
if result:
return self._to_tool_model(result, db=db)
return await self._to_tool_model(result, db=db)
else:
return None
except Exception as e:
log.exception(f'Error creating a new tool: {e}')
return None
def get_tool_by_id(self, id: str, db: Optional[Session] = None) -> Optional[ToolModel]:
async def get_tool_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[ToolModel]:
try:
with get_db_context(db) as db:
tool = db.get(Tool, id)
return self._to_tool_model(tool, db=db) if tool else None
async with get_async_db_context(db) as db:
tool = await db.get(Tool, id)
return await self._to_tool_model(tool, db=db) if tool else None
except Exception:
return None
def get_tools(self, defer_content: bool = False, db: Optional[Session] = None) -> list[ToolUserModel]:
with get_db_context(db) as db:
query = db.query(Tool).order_by(Tool.updated_at.desc())
async def get_tools(self, defer_content: bool = False, db: Optional[AsyncSession] = None) -> list[ToolUserModel]:
async with get_async_db_context(db) as db:
stmt = select(Tool).order_by(Tool.updated_at.desc())
if defer_content:
query = query.options(defer(Tool.content), defer(Tool.specs))
all_tools = query.all()
stmt = stmt
result = await db.execute(stmt)
all_tools = result.scalars().all()
user_ids = list(set(tool.user_id for tool in all_tools))
tool_ids = [tool.id for tool in all_tools]
users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
users = await Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
users_dict = {user.id: user for user in users}
grants_map = AccessGrants.get_grants_by_resources('tool', tool_ids, db=db)
grants_map = await AccessGrants.get_grants_by_resources('tool', tool_ids, db=db)
tools = []
for tool in all_tools:
@ -170,62 +172,66 @@ class ToolsTable:
tools.append(
ToolUserModel.model_validate(
{
**self._to_tool_model(
**(await self._to_tool_model(
tool,
access_grants=grants_map.get(tool.id, []),
db=db,
).model_dump(),
)).model_dump(),
'user': user.model_dump() if user else None,
}
)
)
return tools
def get_tools_by_user_id(
async def get_tools_by_user_id(
self,
user_id: str,
permission: str = 'write',
defer_content: bool = False,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> list[ToolUserModel]:
tools = self.get_tools(defer_content=defer_content, db=db)
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id, db=db)}
tools = await self.get_tools(defer_content=defer_content, db=db)
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
user_group_ids = {group.id for group in user_groups}
return [
tool
for tool in tools
if tool.user_id == user_id
or AccessGrants.has_access(
result = []
for tool in tools:
if tool.user_id == user_id:
result.append(tool)
elif await AccessGrants.has_access(
user_id=user_id,
resource_type='tool',
resource_id=tool.id,
permission=permission,
user_group_ids=user_group_ids,
db=db,
)
]
):
result.append(tool)
return result
def get_tool_valves_by_id(self, id: str, db: Optional[Session] = None) -> Optional[dict]:
async def get_tool_valves_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[dict]:
try:
with get_db_context(db) as db:
tool = db.get(Tool, id)
async with get_async_db_context(db) as db:
tool = await db.get(Tool, id)
return tool.valves if tool.valves else {}
except Exception as e:
log.exception(f'Error getting tool valves by id {id}')
return None
def update_tool_valves_by_id(self, id: str, valves: dict, db: Optional[Session] = None) -> Optional[ToolValves]:
async def update_tool_valves_by_id(self, id: str, valves: dict, db: Optional[AsyncSession] = None) -> Optional[ToolValves]:
try:
with get_db_context(db) as db:
db.query(Tool).filter_by(id=id).update({'valves': valves, 'updated_at': int(time.time())})
db.commit()
return self.get_tool_by_id(id, db=db)
async with get_async_db_context(db) as db:
await db.execute(
update(Tool).filter_by(id=id).values(valves=valves, updated_at=int(time.time()))
)
await db.commit()
return await self.get_tool_by_id(id, db=db)
except Exception:
return None
def get_user_valves_by_id_and_user_id(self, id: str, user_id: str, db: Optional[Session] = None) -> Optional[dict]:
async def get_user_valves_by_id_and_user_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> Optional[dict]:
try:
user = Users.get_user_by_id(user_id, db=db)
user = await Users.get_user_by_id(user_id, db=db)
user_settings = user.settings.model_dump() if user.settings else {}
# Check if user has "tools" and "valves" settings
@ -239,11 +245,11 @@ class ToolsTable:
log.exception(f'Error getting user values by id {id} and user_id {user_id}: {e}')
return None
def update_user_valves_by_id_and_user_id(
self, id: str, user_id: str, valves: dict, db: Optional[Session] = None
async def update_user_valves_by_id_and_user_id(
self, id: str, user_id: str, valves: dict, db: Optional[AsyncSession] = None
) -> Optional[dict]:
try:
user = Users.get_user_by_id(user_id, db=db)
user = await Users.get_user_by_id(user_id, db=db)
user_settings = user.settings.model_dump() if user.settings else {}
# Check if user has "tools" and "valves" settings
@ -255,34 +261,36 @@ class ToolsTable:
user_settings['tools']['valves'][id] = valves
# Update the user settings in the database
Users.update_user_by_id(user_id, {'settings': user_settings}, db=db)
await Users.update_user_by_id(user_id, {'settings': user_settings}, db=db)
return user_settings['tools']['valves'][id]
except Exception as e:
log.exception(f'Error updating user valves by id {id} and user_id {user_id}: {e}')
return None
def update_tool_by_id(self, id: str, updated: dict, db: Optional[Session] = None) -> Optional[ToolModel]:
async def update_tool_by_id(self, id: str, updated: dict, db: Optional[AsyncSession] = None) -> Optional[ToolModel]:
try:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
access_grants = updated.pop('access_grants', None)
db.query(Tool).filter_by(id=id).update({**updated, 'updated_at': int(time.time())})
db.commit()
await db.execute(
update(Tool).filter_by(id=id).values(**updated, updated_at=int(time.time()))
)
await db.commit()
if access_grants is not None:
AccessGrants.set_access_grants('tool', id, access_grants, db=db)
await AccessGrants.set_access_grants('tool', id, access_grants, db=db)
tool = db.query(Tool).get(id)
db.refresh(tool)
return self._to_tool_model(tool, db=db)
tool = await db.get(Tool, id)
await db.refresh(tool)
return await self._to_tool_model(tool, db=db)
except Exception:
return None
def delete_tool_by_id(self, id: str, db: Optional[Session] = None) -> bool:
async def delete_tool_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
try:
with get_db_context(db) as db:
AccessGrants.revoke_all_access('tool', id, db=db)
db.query(Tool).filter_by(id=id).delete()
db.commit()
async with get_async_db_context(db) as db:
await AccessGrants.revoke_all_access('tool', id, db=db)
await db.execute(delete(Tool).filter_by(id=id))
await db.commit()
return True
except Exception:

View file

@ -1,20 +1,15 @@
import time
from typing import Optional
from sqlalchemy.orm import Session, defer
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
from sqlalchemy import select, delete, update, func, or_, case, exists
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import Base, JSONField, get_async_db_context
from open_webui.env import DATABASE_USER_ACTIVE_STATUS_UPDATE_INTERVAL
from open_webui.models.chats import Chats
from open_webui.models.groups import Groups, GroupMember
from open_webui.models.channels import ChannelMember
from open_webui.utils.misc import throttle
from open_webui.utils.validate import validate_profile_image_url
from pydantic import BaseModel, ConfigDict, field_validator, model_validator
from sqlalchemy import (
BigInteger,
@ -24,11 +19,8 @@ from sqlalchemy import (
Boolean,
Text,
Date,
exists,
select,
cast,
)
from sqlalchemy import or_, case, func
from sqlalchemy.dialects.postgresql import JSONB
import datetime
@ -39,13 +31,11 @@ import datetime
# daily bread of every session. Let none go hungry.
####################
class UserSettings(BaseModel):
ui: Optional[dict] = {}
model_config = ConfigDict(extra='allow')
pass
class User(Base):
__tablename__ = 'user'
@ -79,7 +69,6 @@ class User(Base):
updated_at = Column(BigInteger)
created_at = Column(BigInteger)
class UserModel(BaseModel):
id: str
@ -120,13 +109,11 @@ class UserModel(BaseModel):
self.profile_image_url = f'/api/v1/users/{self.id}/profile/image'
return self
class UserStatusModel(UserModel):
is_active: bool = False
model_config = ConfigDict(from_attributes=True)
class ApiKey(Base):
__tablename__ = 'api_key'
@ -139,7 +126,6 @@ class ApiKey(Base):
created_at = Column(BigInteger, nullable=False)
updated_at = Column(BigInteger, nullable=False)
class ApiKeyModel(BaseModel):
id: str
user_id: str
@ -152,12 +138,10 @@ class ApiKeyModel(BaseModel):
model_config = ConfigDict(from_attributes=True)
####################
# Forms
####################
class UpdateProfileForm(BaseModel):
profile_image_url: str
name: str
@ -170,31 +154,25 @@ class UpdateProfileForm(BaseModel):
def check_profile_image_url(cls, v: str) -> str:
return validate_profile_image_url(v)
class UserGroupIdsModel(UserModel):
group_ids: list[str] = []
class UserModelResponse(UserModel):
model_config = ConfigDict(extra='allow')
class UserListResponse(BaseModel):
users: list[UserModelResponse]
total: int
class UserGroupIdsListResponse(BaseModel):
users: list[UserGroupIdsModel]
total: int
class UserStatus(BaseModel):
status_emoji: Optional[str] = None
status_message: Optional[str] = None
status_expires_at: Optional[int] = None
class UserInfoResponse(UserStatus):
id: str
name: str
@ -204,48 +182,39 @@ class UserInfoResponse(UserStatus):
groups: Optional[list] = []
is_active: bool = False
class UserIdNameResponse(BaseModel):
id: str
name: str
class UserIdNameStatusResponse(UserStatus):
id: str
name: str
is_active: Optional[bool] = None
class UserInfoListResponse(BaseModel):
users: list[UserInfoResponse]
total: int
class UserIdNameListResponse(BaseModel):
users: list[UserIdNameResponse]
total: int
class UserNameResponse(BaseModel):
id: str
name: str
role: str
class UserResponse(UserNameResponse):
email: str
class UserProfileImageResponse(UserNameResponse):
email: str
profile_image_url: str
class UserRoleUpdateForm(BaseModel):
id: str
role: str
class UserUpdateForm(BaseModel):
role: str
name: str
@ -258,9 +227,8 @@ class UserUpdateForm(BaseModel):
def check_profile_image_url(cls, v: str) -> str:
return validate_profile_image_url(v)
class UsersTable:
def insert_new_user(
async def insert_new_user(
self,
id: str,
name: str,
@ -269,9 +237,9 @@ class UsersTable:
role: str = 'pending',
username: Optional[str] = None,
oauth: Optional[dict] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[UserModel]:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
user = UserModel(
**{
'id': id,
@ -288,87 +256,98 @@ class UsersTable:
)
result = User(**user.model_dump())
db.add(result)
db.commit()
db.refresh(result)
await db.commit()
await db.refresh(result)
if result:
return user
else:
return None
def get_user_by_id(self, id: str, db: Optional[Session] = None) -> Optional[UserModel]:
async def get_user_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[UserModel]:
try:
with get_db_context(db) as db:
user = db.query(User).filter_by(id=id).first()
return UserModel.model_validate(user)
except Exception:
return None
def get_user_by_api_key(self, api_key: str, db: Optional[Session] = None) -> Optional[UserModel]:
try:
with get_db_context(db) as db:
user = db.query(User).join(ApiKey, User.id == ApiKey.user_id).filter(ApiKey.key == api_key).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter_by(id=id))
user = result.scalars().first()
return UserModel.model_validate(user) if user else None
except Exception:
return None
def get_user_by_email(self, email: str, db: Optional[Session] = None) -> Optional[UserModel]:
async def get_user_by_api_key(self, api_key: str, db: Optional[AsyncSession] = None) -> Optional[UserModel]:
try:
with get_db_context(db) as db:
user = db.query(User).filter(func.lower(User.email) == email.lower()).first()
async with get_async_db_context(db) as db:
result = await db.execute(
select(User).join(ApiKey, User.id == ApiKey.user_id).filter(ApiKey.key == api_key)
)
user = result.scalars().first()
return UserModel.model_validate(user) if user else None
except Exception:
return None
def get_user_by_oauth_sub(self, provider: str, sub: str, db: Optional[Session] = None) -> Optional[UserModel]:
async def get_user_by_email(self, email: str, db: Optional[AsyncSession] = None) -> Optional[UserModel]:
try:
with get_db_context(db) as db: # type: Session
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter(func.lower(User.email) == email.lower()))
user = result.scalars().first()
return UserModel.model_validate(user) if user else None
except Exception:
return None
async def get_user_by_oauth_sub(self, provider: str, sub: str, db: Optional[AsyncSession] = None) -> Optional[UserModel]:
try:
async with get_async_db_context(db) as db:
dialect_name = db.bind.dialect.name
query = db.query(User)
stmt = select(User)
if dialect_name == 'sqlite':
query = query.filter(User.oauth.contains({provider: {'sub': sub}}))
stmt = stmt.filter(User.oauth.contains({provider: {'sub': sub}}))
elif dialect_name == 'postgresql':
query = query.filter(User.oauth[provider].cast(JSONB)['sub'].astext == sub)
stmt = stmt.filter(User.oauth[provider].cast(JSONB)['sub'].astext == sub)
user = query.first()
result = await db.execute(stmt)
user = result.scalars().first()
return UserModel.model_validate(user) if user else None
except Exception as e:
# You may want to log the exception here
return None
def get_user_by_scim_external_id(
self, provider: str, external_id: str, db: Optional[Session] = None
async def get_user_by_scim_external_id(
self, provider: str, external_id: str, db: Optional[AsyncSession] = None
) -> Optional[UserModel]:
try:
with get_db_context(db) as db: # type: Session
async with get_async_db_context(db) as db:
dialect_name = db.bind.dialect.name
query = db.query(User)
stmt = select(User)
if dialect_name == 'sqlite':
query = query.filter(User.scim.contains({provider: {'external_id': external_id}}))
stmt = stmt.filter(User.scim.contains({provider: {'external_id': external_id}}))
elif dialect_name == 'postgresql':
query = query.filter(User.scim[provider].cast(JSONB)['external_id'].astext == external_id)
stmt = stmt.filter(User.scim[provider].cast(JSONB)['external_id'].astext == external_id)
user = query.first()
result = await db.execute(stmt)
user = result.scalars().first()
return UserModel.model_validate(user) if user else None
except Exception:
return None
def get_users(
async def get_users(
self,
filter: Optional[dict] = None,
skip: Optional[int] = None,
limit: Optional[int] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> dict:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
# Import here to avoid circular imports
from open_webui.models.groups import GroupMember
from open_webui.models.channels import ChannelMember
# Join GroupMember so we can order by group_id when requested
query = db.query(User).options(defer(User.profile_image_url))
stmt = select(User)
if filter:
query_key = filter.get('query')
if query_key:
query = query.filter(
stmt = stmt.filter(
or_(
User.name.ilike(f'%{query_key}%'),
User.email.ilike(f'%{query_key}%'),
@ -377,7 +356,7 @@ class UsersTable:
channel_id = filter.get('channel_id')
if channel_id:
query = query.filter(
stmt = stmt.filter(
exists(
select(ChannelMember.id).where(
ChannelMember.user_id == User.id,
@ -395,10 +374,10 @@ class UsersTable:
return {'users': [], 'total': 0}
if user_ids:
query = query.filter(User.id.in_(user_ids))
stmt = stmt.filter(User.id.in_(user_ids))
if group_ids:
query = query.filter(
stmt = stmt.filter(
exists(
select(GroupMember.id).where(
GroupMember.user_id == User.id,
@ -413,9 +392,9 @@ class UsersTable:
exclude_roles = [role[1:] for role in roles if role.startswith('!')]
if include_roles:
query = query.filter(User.role.in_(include_roles))
stmt = stmt.filter(User.role.in_(include_roles))
if exclude_roles:
query = query.filter(~User.role.in_(exclude_roles))
stmt = stmt.filter(~User.role.in_(exclude_roles))
order_by = filter.get('order_by')
direction = filter.get('direction')
@ -435,99 +414,111 @@ class UsersTable:
group_sort = case((membership_exists, 1), else_=0)
if direction == 'asc':
query = query.order_by(group_sort.asc(), User.name.asc())
stmt = stmt.order_by(group_sort.asc(), User.name.asc())
else:
query = query.order_by(group_sort.desc(), User.name.asc())
stmt = stmt.order_by(group_sort.desc(), User.name.asc())
elif order_by == 'name':
if direction == 'asc':
query = query.order_by(User.name.asc())
stmt = stmt.order_by(User.name.asc())
else:
query = query.order_by(User.name.desc())
stmt = stmt.order_by(User.name.desc())
elif order_by == 'email':
if direction == 'asc':
query = query.order_by(User.email.asc())
stmt = stmt.order_by(User.email.asc())
else:
query = query.order_by(User.email.desc())
stmt = stmt.order_by(User.email.desc())
elif order_by == 'created_at':
if direction == 'asc':
query = query.order_by(User.created_at.asc())
stmt = stmt.order_by(User.created_at.asc())
else:
query = query.order_by(User.created_at.desc())
stmt = stmt.order_by(User.created_at.desc())
elif order_by == 'last_active_at':
if direction == 'asc':
query = query.order_by(User.last_active_at.asc())
stmt = stmt.order_by(User.last_active_at.asc())
else:
query = query.order_by(User.last_active_at.desc())
stmt = stmt.order_by(User.last_active_at.desc())
elif order_by == 'updated_at':
if direction == 'asc':
query = query.order_by(User.updated_at.asc())
stmt = stmt.order_by(User.updated_at.asc())
else:
query = query.order_by(User.updated_at.desc())
stmt = stmt.order_by(User.updated_at.desc())
elif order_by == 'role':
if direction == 'asc':
query = query.order_by(User.role.asc())
stmt = stmt.order_by(User.role.asc())
else:
query = query.order_by(User.role.desc())
stmt = stmt.order_by(User.role.desc())
else:
query = query.order_by(User.created_at.desc())
stmt = stmt.order_by(User.created_at.desc())
# Count BEFORE pagination
total = query.count()
count_result = await db.execute(
select(func.count()).select_from(stmt.subquery())
)
total = count_result.scalar()
# correct pagination logic
if skip is not None:
query = query.offset(skip)
stmt = stmt.offset(skip)
if limit is not None:
query = query.limit(limit)
stmt = stmt.limit(limit)
users = query.all()
result = await db.execute(stmt)
users = result.scalars().all()
return {
'users': [UserModel.model_validate(user) for user in users],
'total': total,
}
def get_users_by_group_id(self, group_id: str, db: Optional[Session] = None) -> list[UserModel]:
with get_db_context(db) as db:
users = (
db.query(User)
.options(defer(User.profile_image_url))
async def get_users_by_group_id(self, group_id: str, db: Optional[AsyncSession] = None) -> list[UserModel]:
async with get_async_db_context(db) as db:
from open_webui.models.groups import GroupMember
result = await db.execute(
select(User)
.join(GroupMember, User.id == GroupMember.user_id)
.filter(GroupMember.group_id == group_id)
.all()
)
users = result.scalars().all()
return [UserModel.model_validate(user) for user in users]
def get_users_by_user_ids(self, user_ids: list[str], db: Optional[Session] = None) -> list[UserStatusModel]:
with get_db_context(db) as db:
users = db.query(User).options(defer(User.profile_image_url)).filter(User.id.in_(user_ids)).all()
async def get_users_by_user_ids(self, user_ids: list[str], db: Optional[AsyncSession] = None) -> list[UserStatusModel]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(User).filter(User.id.in_(user_ids))
)
users = result.scalars().all()
return [UserModel.model_validate(user) for user in users]
def get_num_users(self, db: Optional[Session] = None) -> Optional[int]:
with get_db_context(db) as db:
return db.query(User).count()
async def get_num_users(self, db: Optional[AsyncSession] = None) -> Optional[int]:
async with get_async_db_context(db) as db:
result = await db.execute(select(func.count()).select_from(User))
return result.scalar()
def has_users(self, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
return db.query(db.query(User).exists()).scalar()
async def has_users(self, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
result = await db.execute(select(exists(select(User))))
return result.scalar()
def get_first_user(self, db: Optional[Session] = None) -> UserModel:
async def get_first_user(self, db: Optional[AsyncSession] = None) -> UserModel:
try:
with get_db_context(db) as db:
user = db.query(User).order_by(User.created_at).first()
return UserModel.model_validate(user)
async with get_async_db_context(db) as db:
result = await db.execute(select(User).order_by(User.created_at).limit(1))
user = result.scalars().first()
return UserModel.model_validate(user) if user else None
except Exception:
return None
def get_user_webhook_url_by_id(self, id: str, db: Optional[Session] = None) -> Optional[str]:
async def get_user_webhook_url_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[str]:
try:
with get_db_context(db) as db:
user = db.query(User).filter_by(id=id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter_by(id=id))
user = result.scalars().first()
if user.settings is None:
return None
@ -536,68 +527,73 @@ class UsersTable:
except Exception:
return None
def get_num_users_active_today(self, db: Optional[Session] = None) -> Optional[int]:
with get_db_context(db) as db:
async def get_num_users_active_today(self, db: Optional[AsyncSession] = None) -> Optional[int]:
async with get_async_db_context(db) as db:
current_timestamp = int(datetime.datetime.now().timestamp())
today_midnight_timestamp = current_timestamp - (current_timestamp % 86400)
query = db.query(User).filter(User.last_active_at > today_midnight_timestamp)
return query.count()
result = await db.execute(
select(func.count()).select_from(User).filter(User.last_active_at > today_midnight_timestamp)
)
return result.scalar()
def update_user_role_by_id(self, id: str, role: str, db: Optional[Session] = None) -> Optional[UserModel]:
async def update_user_role_by_id(self, id: str, role: str, db: Optional[AsyncSession] = None) -> Optional[UserModel]:
try:
with get_db_context(db) as db:
user = db.query(User).filter_by(id=id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter_by(id=id))
user = result.scalars().first()
if not user:
return None
user.role = role
db.commit()
db.refresh(user)
await db.commit()
await db.refresh(user)
return UserModel.model_validate(user)
except Exception:
return None
def update_user_status_by_id(
self, id: str, form_data: UserStatus, db: Optional[Session] = None
async def update_user_status_by_id(
self, id: str, form_data: UserStatus, db: Optional[AsyncSession] = None
) -> Optional[UserModel]:
try:
with get_db_context(db) as db:
user = db.query(User).filter_by(id=id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter_by(id=id))
user = result.scalars().first()
if not user:
return None
for key, value in form_data.model_dump(exclude_none=True).items():
setattr(user, key, value)
db.commit()
db.refresh(user)
await db.commit()
await db.refresh(user)
return UserModel.model_validate(user)
except Exception:
return None
def update_user_profile_image_url_by_id(
self, id: str, profile_image_url: str, db: Optional[Session] = None
async def update_user_profile_image_url_by_id(
self, id: str, profile_image_url: str, db: Optional[AsyncSession] = None
) -> Optional[UserModel]:
try:
with get_db_context(db) as db:
user = db.query(User).filter_by(id=id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter_by(id=id))
user = result.scalars().first()
if not user:
return None
user.profile_image_url = profile_image_url
db.commit()
db.refresh(user)
await db.commit()
await db.refresh(user)
return UserModel.model_validate(user)
except Exception:
return None
@throttle(DATABASE_USER_ACTIVE_STATUS_UPDATE_INTERVAL)
def update_last_active_by_id(self, id: str, db: Optional[Session] = None) -> None:
async def update_last_active_by_id(self, id: str, db: Optional[AsyncSession] = None) -> None:
try:
with get_db_context(db) as db:
db.query(User).filter_by(id=id).update({'last_active_at': int(time.time())})
db.commit()
async with get_async_db_context(db) as db:
await db.execute(update(User).filter_by(id=id).values(last_active_at=int(time.time())))
await db.commit()
except Exception:
pass
def update_user_oauth_by_id(
self, id: str, provider: str, sub: str, db: Optional[Session] = None
async def update_user_oauth_by_id(
self, id: str, provider: str, sub: str, db: Optional[AsyncSession] = None
) -> Optional[UserModel]:
"""
Update or insert an OAuth provider/sub pair into the user's oauth JSON field.
@ -608,8 +604,9 @@ class UsersTable:
}
"""
try:
with get_db_context(db) as db:
user = db.query(User).filter_by(id=id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter_by(id=id))
user = result.scalars().first()
if not user:
return None
@ -620,20 +617,20 @@ class UsersTable:
oauth[provider] = {'sub': sub}
# Persist updated JSON
db.query(User).filter_by(id=id).update({'oauth': oauth})
db.commit()
await db.execute(update(User).filter_by(id=id).values(oauth=oauth))
await db.commit()
return UserModel.model_validate(user)
except Exception:
return None
def update_user_scim_by_id(
async def update_user_scim_by_id(
self,
id: str,
provider: str,
external_id: str,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
) -> Optional[UserModel]:
"""
Update or insert a SCIM provider/external_id pair into the user's scim JSON field.
@ -644,41 +641,44 @@ class UsersTable:
}
"""
try:
with get_db_context(db) as db:
user = db.query(User).filter_by(id=id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter_by(id=id))
user = result.scalars().first()
if not user:
return None
scim = user.scim or {}
scim[provider] = {'external_id': external_id}
db.query(User).filter_by(id=id).update({'scim': scim})
db.commit()
await db.execute(update(User).filter_by(id=id).values(scim=scim))
await db.commit()
return UserModel.model_validate(user)
except Exception:
return None
def update_user_by_id(self, id: str, updated: dict, db: Optional[Session] = None) -> Optional[UserModel]:
async def update_user_by_id(self, id: str, updated: dict, db: Optional[AsyncSession] = None) -> Optional[UserModel]:
try:
with get_db_context(db) as db:
user = db.query(User).filter_by(id=id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter_by(id=id))
user = result.scalars().first()
if not user:
return None
for key, value in updated.items():
setattr(user, key, value)
db.commit()
db.refresh(user)
await db.commit()
await db.refresh(user)
return UserModel.model_validate(user)
except Exception as e:
print(e)
return None
def update_user_settings_by_id(self, id: str, updated: dict, db: Optional[Session] = None) -> Optional[UserModel]:
async def update_user_settings_by_id(self, id: str, updated: dict, db: Optional[AsyncSession] = None) -> Optional[UserModel]:
try:
with get_db_context(db) as db:
user = db.query(User).filter_by(id=id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter_by(id=id))
user = result.scalars().first()
if not user:
return None
@ -689,26 +689,30 @@ class UsersTable:
user_settings.update(updated)
db.query(User).filter_by(id=id).update({'settings': user_settings})
db.commit()
await db.execute(update(User).filter_by(id=id).values(settings=user_settings))
await db.commit()
user = db.query(User).filter_by(id=id).first()
result = await db.execute(select(User).filter_by(id=id))
user = result.scalars().first()
return UserModel.model_validate(user)
except Exception:
return None
def delete_user_by_id(self, id: str, db: Optional[Session] = None) -> bool:
async def delete_user_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
try:
from open_webui.models.groups import Groups
from open_webui.models.chats import Chats
# Remove User from Groups
Groups.remove_user_from_all_groups(id)
await Groups.remove_user_from_all_groups(id)
# Delete User Chats
result = Chats.delete_chats_by_user_id(id, db=db)
result = await Chats.delete_chats_by_user_id(id, db=db)
if result:
with get_db_context(db) as db:
async with get_async_db_context(db) as db:
# Delete User
db.query(User).filter_by(id=id).delete()
db.commit()
await db.execute(delete(User).filter_by(id=id))
await db.commit()
return True
else:
@ -716,19 +720,20 @@ class UsersTable:
except Exception:
return False
def get_user_api_key_by_id(self, id: str, db: Optional[Session] = None) -> Optional[str]:
async def get_user_api_key_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[str]:
try:
with get_db_context(db) as db:
api_key = db.query(ApiKey).filter_by(user_id=id).first()
async with get_async_db_context(db) as db:
result = await db.execute(select(ApiKey).filter_by(user_id=id))
api_key = result.scalars().first()
return api_key.key if api_key else None
except Exception:
return None
def update_user_api_key_by_id(self, id: str, api_key: str, db: Optional[Session] = None) -> bool:
async def update_user_api_key_by_id(self, id: str, api_key: str, db: Optional[AsyncSession] = None) -> bool:
try:
with get_db_context(db) as db:
db.query(ApiKey).filter_by(user_id=id).delete()
db.commit()
async with get_async_db_context(db) as db:
await db.execute(delete(ApiKey).filter_by(user_id=id))
await db.commit()
now = int(time.time())
new_api_key = ApiKey(
@ -739,41 +744,45 @@ class UsersTable:
updated_at=now,
)
db.add(new_api_key)
db.commit()
await db.commit()
return True
except Exception:
return False
def delete_user_api_key_by_id(self, id: str, db: Optional[Session] = None) -> bool:
async def delete_user_api_key_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
try:
with get_db_context(db) as db:
db.query(ApiKey).filter_by(user_id=id).delete()
db.commit()
async with get_async_db_context(db) as db:
await db.execute(delete(ApiKey).filter_by(user_id=id))
await db.commit()
return True
except Exception:
return False
def get_valid_user_ids(self, user_ids: list[str], db: Optional[Session] = None) -> list[str]:
with get_db_context(db) as db:
users = db.query(User).filter(User.id.in_(user_ids)).all()
async def get_valid_user_ids(self, user_ids: list[str], db: Optional[AsyncSession] = None) -> list[str]:
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter(User.id.in_(user_ids)))
users = result.scalars().all()
return [user.id for user in users]
def get_super_admin_user(self, db: Optional[Session] = None) -> Optional[UserModel]:
with get_db_context(db) as db:
user = db.query(User).filter_by(role='admin').first()
async def get_super_admin_user(self, db: Optional[AsyncSession] = None) -> Optional[UserModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter_by(role='admin').limit(1))
user = result.scalars().first()
if user:
return UserModel.model_validate(user)
else:
return None
def get_active_user_count(self, db: Optional[Session] = None) -> int:
with get_db_context(db) as db:
async def get_active_user_count(self, db: Optional[AsyncSession] = None) -> int:
async with get_async_db_context(db) as db:
# Consider user active if last_active_at within the last 3 minutes
three_minutes_ago = int(time.time()) - 180
count = db.query(User).filter(User.last_active_at >= three_minutes_ago).count()
return count
result = await db.execute(
select(func.count()).select_from(User).filter(User.last_active_at >= three_minutes_ago)
)
return result.scalar()
@staticmethod
def is_active(user: UserModel) -> bool:
@ -783,14 +792,14 @@ class UsersTable:
return user.last_active_at >= three_minutes_ago
return False
def is_user_active(self, user_id: str, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
user = db.query(User).filter_by(id=user_id).first()
async def is_user_active(self, user_id: str, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter_by(id=user_id))
user = result.scalars().first()
if user and user.last_active_at:
# Consider user active if last_active_at within the last 3 minutes
three_minutes_ago = int(time.time()) - 180
return user.last_active_at >= three_minutes_ago
return False
Users = UsersTable()

View file

@ -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 = []

View file

@ -11,8 +11,8 @@ from open_webui.models.groups import Groups
from open_webui.models.users import Users
from open_webui.models.feedbacks import Feedbacks
from open_webui.utils.auth import get_admin_user
from open_webui.internal.db import get_session
from sqlalchemy.orm import Session
from open_webui.internal.db import get_async_session
from sqlalchemy.ext.asyncio import AsyncSession
log = logging.getLogger(__name__)
@ -59,10 +59,10 @@ async def get_model_analytics(
end_date: Optional[int] = Query(None, description='End timestamp (epoch)'),
group_id: Optional[str] = Query(None, description='Filter by user group ID'),
user=Depends(get_admin_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""Get message counts per model."""
counts = ChatMessages.get_message_count_by_model(start_date=start_date, end_date=end_date, group_id=group_id, db=db)
counts = await ChatMessages.get_message_count_by_model(start_date=start_date, end_date=end_date, group_id=group_id, db=db)
models = [
ModelAnalyticsEntry(model_id=model_id, count=count)
for model_id, count in sorted(counts.items(), key=lambda x: -x[1])
@ -77,17 +77,17 @@ async def get_user_analytics(
group_id: Optional[str] = Query(None, description='Filter by user group ID'),
limit: int = Query(50, description='Max users to return'),
user=Depends(get_admin_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""Get message counts and token usage per user with user info."""
counts = ChatMessages.get_message_count_by_user(start_date=start_date, end_date=end_date, group_id=group_id, db=db)
token_usage = ChatMessages.get_token_usage_by_user(
counts = await ChatMessages.get_message_count_by_user(start_date=start_date, end_date=end_date, group_id=group_id, db=db)
token_usage = await ChatMessages.get_token_usage_by_user(
start_date=start_date, end_date=end_date, group_id=group_id, db=db
)
# Get user info for top users
top_user_ids = [uid for uid, _ in sorted(counts.items(), key=lambda x: -x[1])[:limit]]
user_info = {u.id: u for u in Users.get_users_by_user_ids(top_user_ids, db=db)}
user_info = {u.id: u for u in await Users.get_users_by_user_ids(top_user_ids, db=db)}
users = []
for user_id in top_user_ids:
@ -118,13 +118,13 @@ async def get_messages(
skip: int = Query(0),
limit: int = Query(50, le=100),
user=Depends(get_admin_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""Query messages with filters."""
if chat_id:
return ChatMessages.get_messages_by_chat_id(chat_id=chat_id, db=db)
return await ChatMessages.get_messages_by_chat_id(chat_id=chat_id, db=db)
elif model_id:
return ChatMessages.get_messages_by_model_id(
return await ChatMessages.get_messages_by_model_id(
model_id=model_id,
start_date=start_date,
end_date=end_date,
@ -133,7 +133,7 @@ async def get_messages(
db=db,
)
elif user_id:
return ChatMessages.get_messages_by_user_id(user_id=user_id, skip=skip, limit=limit, db=db)
return await ChatMessages.get_messages_by_user_id(user_id=user_id, skip=skip, limit=limit, db=db)
else:
# Return empty if no filter specified
return []
@ -152,16 +152,16 @@ async def get_summary(
end_date: Optional[int] = Query(None, description='End timestamp (epoch)'),
group_id: Optional[str] = Query(None, description='Filter by user group ID'),
user=Depends(get_admin_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""Get summary statistics for the dashboard."""
model_counts = ChatMessages.get_message_count_by_model(
model_counts = await ChatMessages.get_message_count_by_model(
start_date=start_date, end_date=end_date, group_id=group_id, db=db
)
user_counts = ChatMessages.get_message_count_by_user(
user_counts = await ChatMessages.get_message_count_by_user(
start_date=start_date, end_date=end_date, group_id=group_id, db=db
)
chat_counts = ChatMessages.get_message_count_by_chat(
chat_counts = await ChatMessages.get_message_count_by_chat(
start_date=start_date, end_date=end_date, group_id=group_id, db=db
)
@ -189,13 +189,13 @@ async def get_daily_stats(
group_id: Optional[str] = Query(None, description='Filter by user group ID'),
granularity: str = Query('daily', description="Granularity: 'hourly' or 'daily'"),
user=Depends(get_admin_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""Get message counts grouped by model for time-series chart."""
if granularity == 'hourly':
counts = ChatMessages.get_hourly_message_counts_by_model(start_date=start_date, end_date=end_date, db=db)
counts = await ChatMessages.get_hourly_message_counts_by_model(start_date=start_date, end_date=end_date, db=db)
else:
counts = ChatMessages.get_daily_message_counts_by_model(
counts = await ChatMessages.get_daily_message_counts_by_model(
start_date=start_date, end_date=end_date, group_id=group_id, db=db
)
return DailyStatsResponse(
@ -224,10 +224,10 @@ async def get_token_usage(
end_date: Optional[int] = Query(None),
group_id: Optional[str] = Query(None, description='Filter by user group ID'),
user=Depends(get_admin_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""Get token usage aggregated by model."""
usage = ChatMessages.get_token_usage_by_model(start_date=start_date, end_date=end_date, group_id=group_id, db=db)
usage = await ChatMessages.get_token_usage_by_model(start_date=start_date, end_date=end_date, group_id=group_id, db=db)
models = [
TokenUsageEntry(model_id=model_id, **data)
@ -271,12 +271,12 @@ async def get_model_chats(
skip: int = Query(0),
limit: int = Query(50, le=100),
user=Depends(get_admin_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""Get chats that used a specific model, with preview and feedback info."""
# Get chat IDs that used this model
chat_ids = ChatMessages.get_chat_ids_by_model_id(
chat_ids = await ChatMessages.get_chat_ids_by_model_id(
model_id=model_id,
start_date=start_date,
end_date=end_date,
@ -291,7 +291,7 @@ async def get_model_chats(
# Get chat details from messages only
chats_data = []
for chat_id in chat_ids:
messages = ChatMessages.get_messages_by_chat_id(chat_id, db=db)
messages = await ChatMessages.get_messages_by_chat_id(chat_id, db=db)
if not messages:
continue
@ -312,7 +312,7 @@ async def get_model_chats(
# Get user info
user_name = None
if user_id:
user_info = Users.get_user_by_id(user_id, db=db)
user_info = await Users.get_user_by_id(user_id, db=db)
user_name = user_info.name if user_info else None
# Timestamps from messages
@ -357,12 +357,12 @@ async def get_model_overview(
model_id: str,
days: int = Query(30, description='Number of days of history (0 for all)'),
user=Depends(get_admin_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""Get model overview with feedback history and chat tags."""
# Get chat IDs that used this model
chat_ids = ChatMessages.get_chat_ids_by_model_id(
chat_ids = await ChatMessages.get_chat_ids_by_model_id(
model_id=model_id,
start_date=None,
end_date=None,
@ -381,7 +381,7 @@ async def get_model_overview(
start_dt = now - timedelta(days=days)
for chat_id in chat_ids:
feedbacks = Feedbacks.get_feedbacks_by_chat_id(chat_id, db=db)
feedbacks = await Feedbacks.get_feedbacks_by_chat_id(chat_id, db=db)
for fb in feedbacks:
if fb.data and 'rating' in fb.data:
rating = fb.data['rating']
@ -425,7 +425,7 @@ async def get_model_overview(
# Get chat tags
tag_counts: dict[str, int] = defaultdict(int)
for chat_id in chat_ids:
chat = Chats.get_chat_by_id(chat_id, db=db)
chat = await Chats.get_chat_by_id(chat_id, db=db)
if chat and chat.meta:
for tag in chat.meta.get('tags', []):
tag_counts[tag] += 1

View file

@ -330,7 +330,7 @@ async def speech(request: Request, user=Depends(get_verified_user)):
detail=ERROR_MESSAGES.NOT_FOUND,
)
if user.role != 'admin' and not has_permission(user.id, 'chat.tts', request.app.state.config.USER_PERMISSIONS):
if user.role != 'admin' and not await has_permission(user.id, 'chat.tts', request.app.state.config.USER_PERMISSIONS):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
@ -660,7 +660,7 @@ def transcription_handler(request, file_path, metadata, user=None):
data = {'text': transcript.strip()}
# save the transcript to a json file
transcript_file = f'{file_dir}/{id}.json'
transcript_file = os.path.join(file_dir, f'{id}.json')
with open(transcript_file, 'w') as f:
json.dump(data, f)
@ -698,7 +698,7 @@ def transcription_handler(request, file_path, metadata, user=None):
data = r.json()
# save the transcript to a json file
transcript_file = f'{file_dir}/{id}.json'
transcript_file = os.path.join(file_dir, f'{id}.json')
with open(transcript_file, 'w') as f:
json.dump(data, f)
@ -767,7 +767,7 @@ def transcription_handler(request, file_path, metadata, user=None):
data = {'text': transcript.strip()}
# Save transcript
transcript_file = f'{file_dir}/{id}.json'
transcript_file = os.path.join(file_dir, f'{id}.json')
with open(transcript_file, 'w') as f:
json.dump(data, f)
@ -874,7 +874,7 @@ def transcription_handler(request, file_path, metadata, user=None):
data = {'text': transcript}
# Save transcript to json file (consistent with other providers)
transcript_file = f'{file_dir}/{id}.json'
transcript_file = os.path.join(file_dir, f'{id}.json')
with open(transcript_file, 'w') as f:
json.dump(data, f)
@ -1059,7 +1059,7 @@ def transcription_handler(request, file_path, metadata, user=None):
data = {'text': transcript}
# Save transcript to json file (consistent with other providers)
transcript_file = f'{file_dir}/{id}.json'
transcript_file = os.path.join(file_dir, f'{id}.json')
with open(transcript_file, 'w') as f:
json.dump(data, f)
@ -1208,13 +1208,13 @@ def split_audio(file_path, max_bytes, format='mp3', bitrate='32k'):
@router.post('/transcriptions')
def transcription(
async def transcription(
request: Request,
file: UploadFile = File(...),
language: Optional[str] = Form(None),
user=Depends(get_verified_user),
):
if user.role != 'admin' and not has_permission(user.id, 'chat.stt', request.app.state.config.USER_PERMISSIONS):
if user.role != 'admin' and not await has_permission(user.id, 'chat.stt', request.app.state.config.USER_PERMISSIONS):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
@ -1237,9 +1237,9 @@ def transcription(
filename = f'{id}.{ext}'
contents = file.file.read()
file_dir = f'{CACHE_DIR}/audio/transcriptions'
file_dir = os.path.join(CACHE_DIR, 'audio', 'transcriptions')
os.makedirs(file_dir, exist_ok=True)
file_path = f'{file_dir}/{filename}'
file_path = os.path.join(file_dir, filename)
# Defense-in-depth: ensure resolved path stays within intended directory
if not os.path.realpath(file_path).startswith(os.path.realpath(file_dir)):

View file

@ -70,8 +70,8 @@ from open_webui.utils.auth import (
get_password_hash,
get_http_authorization_cred,
)
from open_webui.internal.db import get_session
from sqlalchemy.orm import Session
from open_webui.internal.db import get_async_session
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.utils.webhook import post_webhook
from open_webui.utils.access_control import get_permissions, has_permission
from open_webui.utils.groups import apply_default_group_assignment
@ -96,7 +96,7 @@ log = logging.getLogger(__name__)
signin_rate_limiter = RateLimiter(redis_client=get_redis_client(), limit=5 * 3, window=60 * 3)
def create_session_response(request: Request, user, db, response: Response = None, set_cookie: bool = False) -> dict:
async def create_session_response(request: Request, user, db, response: Response = None, set_cookie: bool = False) -> dict:
"""
Create JWT token and build session response for a user.
Shared helper for signin, signup, ldap_auth, add_user, and token_exchange endpoints.
@ -131,7 +131,7 @@ def create_session_response(request: Request, user, db, response: Response = Non
**({'max_age': max_age} if max_age is not None else {}),
)
user_permissions = get_permissions(user.id, request.app.state.config.USER_PERMISSIONS, db=db)
user_permissions = await get_permissions(user.id, request.app.state.config.USER_PERMISSIONS, db=db)
return {
'token': token,
@ -167,7 +167,7 @@ async def get_session_user(
request: Request,
response: Response,
user=Depends(get_current_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
auth_header = request.headers.get('Authorization')
auth_token = get_http_authorization_cred(auth_header)
@ -197,7 +197,7 @@ async def get_session_user(
**({'max_age': max_age} if max_age is not None else {}),
)
user_permissions = get_permissions(user.id, request.app.state.config.USER_PERMISSIONS, db=db)
user_permissions = await get_permissions(user.id, request.app.state.config.USER_PERMISSIONS, db=db)
return {
'token': token,
@ -227,10 +227,10 @@ async def get_session_user(
async def update_profile(
form_data: UpdateProfileForm,
session_user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if session_user:
user = Users.update_user_by_id(
user = await Users.update_user_by_id(
session_user.id,
form_data.model_dump(),
db=db,
@ -256,10 +256,10 @@ class UpdateTimezoneForm(BaseModel):
async def update_timezone(
form_data: UpdateTimezoneForm,
session_user=Depends(get_current_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if session_user:
Users.update_user_by_id(
await Users.update_user_by_id(
session_user.id,
{'timezone': form_data.timezone},
db=db,
@ -278,12 +278,12 @@ async def update_timezone(
async def update_password(
form_data: UpdatePasswordForm,
session_user=Depends(get_current_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if WEBUI_AUTH_TRUSTED_EMAIL_HEADER:
raise HTTPException(400, detail=ERROR_MESSAGES.ACTION_PROHIBITED)
if session_user:
user = Auths.authenticate_user(
user = await Auths.authenticate_user(
session_user.email,
lambda pw: verify_password(form_data.password, pw),
db=db,
@ -295,7 +295,7 @@ async def update_password(
except Exception as e:
raise HTTPException(400, detail=str(e))
hashed = get_password_hash(form_data.new_password)
return Auths.update_user_password_by_id(user.id, hashed, db=db)
return await Auths.update_user_password_by_id(user.id, hashed, db=db)
else:
raise HTTPException(400, detail=ERROR_MESSAGES.INCORRECT_PASSWORD)
else:
@ -310,7 +310,7 @@ async def ldap_auth(
request: Request,
response: Response,
form_data: LdapForm,
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
# Security checks FIRST - before loading any config
if not request.app.state.config.ENABLE_LDAP:
@ -476,23 +476,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 +510,19 @@ async def ldap_auth(
log.error(f'LDAP user creation error: {str(err)}')
raise HTTPException(500, detail='Internal error occurred during LDAP user creation.')
user = Auths.authenticate_user_by_email(email, db=db)
user = await Auths.authenticate_user_by_email(email, db=db)
if user:
if ENABLE_LDAP_GROUP_MANAGEMENT and user_groups:
if ENABLE_LDAP_GROUP_CREATION:
Groups.create_groups_by_group_names(user.id, user_groups, db=db)
await Groups.create_groups_by_group_names(user.id, user_groups, db=db)
try:
Groups.sync_groups_by_group_names(user.id, user_groups, db=db)
await Groups.sync_groups_by_group_names(user.id, user_groups, db=db)
log.info(f'Successfully synced groups for user {user.id}: {user_groups}')
except Exception as e:
log.error(f'Failed to sync groups for user {user.id}: {e}')
return create_session_response(request, user, db, response, set_cookie=True)
return await create_session_response(request, user, db, response, set_cookie=True)
else:
raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED)
else:
@ -536,7 +542,7 @@ async def signin(
request: Request,
response: Response,
form_data: SigninForm,
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if not ENABLE_PASSWORD_AUTH:
raise HTTPException(
@ -558,7 +564,7 @@ async def signin(
except Exception as e:
pass
if not Users.get_user_by_email(email.lower(), db=db):
if not await Users.get_user_by_email(email.lower(), db=db):
await signup_handler(
request,
email,
@ -567,20 +573,20 @@ async def signin(
db=db,
)
user = Auths.authenticate_user_by_email(email, db=db)
user = await Auths.authenticate_user_by_email(email, db=db)
if user:
if WEBUI_AUTH_TRUSTED_GROUPS_HEADER:
group_names = request.headers.get(WEBUI_AUTH_TRUSTED_GROUPS_HEADER, '').split(',')
group_names = [name.strip() for name in group_names if name.strip()]
if group_names:
Groups.sync_groups_by_group_names(user.id, group_names, db=db)
await Groups.sync_groups_by_group_names(user.id, group_names, db=db)
if WEBUI_AUTH_TRUSTED_ROLE_HEADER:
trusted_role = request.headers.get(WEBUI_AUTH_TRUSTED_ROLE_HEADER, '').lower().strip()
if trusted_role in {'admin', 'user', 'pending'}:
if user.role != trusted_role:
Users.update_user_role_by_id(user.id, trusted_role, db=db)
await Users.update_user_role_by_id(user.id, trusted_role, db=db)
elif trusted_role:
log.warning(f'Ignoring invalid trusted role header value: {trusted_role}')
@ -588,14 +594,14 @@ async def signin(
admin_email = 'admin@localhost'
admin_password = 'admin'
if Users.get_user_by_email(admin_email.lower(), db=db):
user = Auths.authenticate_user(
if await Users.get_user_by_email(admin_email.lower(), db=db):
user = await Auths.authenticate_user(
admin_email.lower(),
lambda pw: verify_password(admin_password, pw),
db=db,
)
else:
if Users.has_users(db=db):
if await Users.has_users(db=db):
raise HTTPException(400, detail=ERROR_MESSAGES.EXISTING_USERS)
await signup_handler(
@ -606,7 +612,7 @@ async def signin(
db=db,
)
user = Auths.authenticate_user(
user = await Auths.authenticate_user(
admin_email.lower(),
lambda pw: verify_password(admin_password, pw),
db=db,
@ -627,14 +633,14 @@ async def signin(
# decode safely — ignore incomplete UTF-8 sequences
form_data.password = password_bytes.decode('utf-8', errors='ignore')
user = Auths.authenticate_user(
user = await Auths.authenticate_user(
form_data.email.lower(),
lambda pw: verify_password(form_data.password, pw),
db=db,
)
if user:
return create_session_response(request, user, db, response, set_cookie=True)
return await create_session_response(request, user, db, response, set_cookie=True)
else:
raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED)
@ -651,7 +657,7 @@ async def signup_handler(
name: str,
profile_image_url: str = '/user.png',
*,
db: Session,
db: AsyncSession,
) -> UserModel:
"""
Core user-creation logic shared by the signup endpoint and
@ -665,7 +671,7 @@ async def signup_handler(
# first-user registration can all see an empty table and each get admin.
hashed = get_password_hash(password)
user = Auths.insert_new_auth(
user = await Auths.insert_new_auth(
email=email.lower(),
password=hashed,
name=name,
@ -678,9 +684,9 @@ async def signup_handler(
# Atomically check if this is the only user *after* the insert.
# Only the single user present at this point should become admin.
if Users.get_num_users(db=db) == 1:
Users.update_user_role_by_id(user.id, 'admin', db=db)
user = Users.get_user_by_id(user.id, db=db)
if await Users.get_num_users(db=db) == 1:
await Users.update_user_role_by_id(user.id, 'admin', db=db)
user = await Users.get_user_by_id(user.id, db=db)
request.app.state.config.ENABLE_SIGNUP = False
if request.app.state.config.WEBHOOK_URL:
@ -695,7 +701,7 @@ async def signup_handler(
},
)
apply_default_group_assignment(
await apply_default_group_assignment(
request.app.state.config.DEFAULT_GROUP_ID,
user.id,
db=db,
@ -709,9 +715,9 @@ async def signup(
request: Request,
response: Response,
form_data: SignupForm,
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
has_users = Users.has_users(db=db)
has_users = await Users.has_users(db=db)
if WEBUI_AUTH:
if not request.app.state.config.ENABLE_SIGNUP or not request.app.state.config.ENABLE_LOGIN_FORM:
@ -724,7 +730,7 @@ async def signup(
if not validate_email_format(form_data.email.lower()):
raise HTTPException(status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.INVALID_EMAIL_FORMAT)
if Users.get_user_by_email(form_data.email.lower(), db=db):
if await Users.get_user_by_email(form_data.email.lower(), db=db):
raise HTTPException(400, detail=ERROR_MESSAGES.EMAIL_TAKEN)
try:
@ -741,7 +747,7 @@ async def signup(
form_data.profile_image_url,
db=db,
)
return create_session_response(request, user, db, response, set_cookie=True)
return await create_session_response(request, user, db, response, set_cookie=True)
except HTTPException:
raise
except Exception as err:
@ -750,7 +756,7 @@ async def signup(
@router.get('/signout')
async def signout(request: Request, response: Response, db: Session = Depends(get_session)):
async def signout(request: Request, response: Response, db: AsyncSession = Depends(get_async_session)):
# get auth token from headers or cookies
token = None
auth_header = request.headers.get('Authorization')
@ -771,7 +777,7 @@ async def signout(request: Request, response: Response, db: Session = Depends(ge
if oauth_session_id:
response.delete_cookie('oauth_session_id')
session = OAuthSessions.get_session_by_id(oauth_session_id, db=db)
session = await OAuthSessions.get_session_by_id(oauth_session_id, db=db)
# If a custom end_session_endpoint is configured (e.g. AWS Cognito), redirect
# there directly instead of attempting OIDC discovery.
@ -846,12 +852,12 @@ async def add_user(
request: Request,
form_data: AddUserForm,
user=Depends(get_admin_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if not validate_email_format(form_data.email.lower()):
raise HTTPException(status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.INVALID_EMAIL_FORMAT)
if Users.get_user_by_email(form_data.email.lower(), db=db):
if await Users.get_user_by_email(form_data.email.lower(), db=db):
raise HTTPException(400, detail=ERROR_MESSAGES.EMAIL_TAKEN)
try:
@ -861,7 +867,7 @@ async def add_user(
raise HTTPException(400, detail=str(e))
hashed = get_password_hash(form_data.password)
user = Auths.insert_new_auth(
user = await Auths.insert_new_auth(
form_data.email.lower(),
hashed,
form_data.name,
@ -871,7 +877,7 @@ async def add_user(
)
if user:
apply_default_group_assignment(
await apply_default_group_assignment(
request.app.state.config.DEFAULT_GROUP_ID,
user.id,
db=db,
@ -903,7 +909,7 @@ async def add_user(
@router.get('/admin/details')
async def get_admin_details(request: Request, user=Depends(get_current_user), db: Session = Depends(get_session)):
async def get_admin_details(request: Request, user=Depends(get_current_user), db: AsyncSession = Depends(get_async_session)):
if request.app.state.config.SHOW_ADMIN_DETAILS:
admin_email = request.app.state.config.ADMIN_EMAIL
admin_name = None
@ -911,11 +917,11 @@ async def get_admin_details(request: Request, user=Depends(get_current_user), db
log.info(f'Admin details - Email: {admin_email}, Name: {admin_name}')
if admin_email:
admin = Users.get_user_by_email(admin_email, db=db)
admin = await Users.get_user_by_email(admin_email, db=db)
if admin:
admin_name = admin.name
else:
admin = Users.get_first_user(db=db)
admin = await Users.get_first_user(db=db)
if admin:
admin_email = admin.email
admin_name = admin.name
@ -1167,10 +1173,10 @@ async def update_ldap_config(request: Request, form_data: LdapConfigForm, user=D
# create api key
@router.post('/api_key', response_model=ApiKey)
async def generate_api_key(request: Request, user=Depends(get_current_user), db: Session = Depends(get_session)):
async def generate_api_key(request: Request, user=Depends(get_current_user), db: AsyncSession = Depends(get_async_session)):
if not request.app.state.config.ENABLE_API_KEYS or (
user.role != 'admin'
and not has_permission(user.id, 'features.api_keys', request.app.state.config.USER_PERMISSIONS)
and not await has_permission(user.id, 'features.api_keys', request.app.state.config.USER_PERMISSIONS)
):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
@ -1178,7 +1184,7 @@ async def generate_api_key(request: Request, user=Depends(get_current_user), db:
)
api_key = create_api_key()
success = Users.update_user_api_key_by_id(user.id, api_key, db=db)
success = await Users.update_user_api_key_by_id(user.id, api_key, db=db)
if success:
return {
@ -1190,14 +1196,14 @@ async def generate_api_key(request: Request, user=Depends(get_current_user), db:
# delete api key
@router.delete('/api_key', response_model=bool)
async def delete_api_key(user=Depends(get_current_user), db: Session = Depends(get_session)):
return Users.delete_user_api_key_by_id(user.id, db=db)
async def delete_api_key(user=Depends(get_current_user), db: AsyncSession = Depends(get_async_session)):
return await Users.delete_user_api_key_by_id(user.id, db=db)
# get api key
@router.get('/api_key', response_model=ApiKey)
async def get_api_key(user=Depends(get_current_user), db: Session = Depends(get_session)):
api_key = Users.get_user_api_key_by_id(user.id, db=db)
async def get_api_key(user=Depends(get_current_user), db: AsyncSession = Depends(get_async_session)):
api_key = await Users.get_user_api_key_by_id(user.id, db=db)
if api_key:
return {
'api_key': api_key,
@ -1221,7 +1227,7 @@ async def token_exchange(
response: Response,
provider: str,
form_data: TokenExchangeForm,
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""
Exchange an external OAuth provider token for an OpenWebUI JWT.
@ -1290,14 +1296,14 @@ async def token_exchange(
email = email.lower()
# Try to find the user by OAuth sub
user = Users.get_user_by_oauth_sub(provider, sub, db=db)
user = await Users.get_user_by_oauth_sub(provider, sub, db=db)
if not user and OAUTH_MERGE_ACCOUNTS_BY_EMAIL.value:
# Try to find by email if merge is enabled
user = Users.get_user_by_email(email, db=db)
user = await Users.get_user_by_email(email, db=db)
if user:
# Link the OAuth sub to this user
Users.update_user_oauth_by_id(user.id, provider, sub, db=db)
await Users.update_user_oauth_by_id(user.id, provider, sub, db=db)
if not user:
raise HTTPException(
@ -1305,4 +1311,4 @@ async def token_exchange(
detail='User not found. Please sign in via the web interface first.',
)
return create_session_response(request, user, db)
return await create_session_response(request, user, db)

View file

@ -3,7 +3,7 @@ import logging
from typing import Optional
from fastapi import APIRouter, Depends, HTTPException, Request, status
from sqlalchemy.orm import Session
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.models.automations import (
Automations,
@ -23,7 +23,7 @@ from open_webui.utils.automations import (
)
from open_webui.utils.auth import get_verified_user, get_admin_user
from open_webui.utils.access_control import has_permission
from open_webui.internal.db import get_session
from open_webui.internal.db import get_async_session
from open_webui.constants import ERROR_MESSAGES
log = logging.getLogger(__name__)
@ -38,8 +38,8 @@ PAGE_ITEM_COUNT = 30
############################
def check_automations_permission(request, user):
if user.role != 'admin' and not has_permission(
async def check_automations_permission(request, user):
if user.role != 'admin' and not await has_permission(
user.id, 'features.automations', request.app.state.config.USER_PERMISSIONS
):
raise HTTPException(
@ -61,7 +61,7 @@ def check_automation_access(automation, user):
)
def check_automation_limits(request, user, rrule_str: str, db, is_create: bool = False):
async def check_automation_limits(request, user, rrule_str: str, db, is_create: bool = False):
"""Enforce global automation limits. Admins bypass all checks."""
if user.role == 'admin':
return
@ -71,7 +71,7 @@ def check_automation_limits(request, user, rrule_str: str, db, is_create: bool =
max_count = request.app.state.config.AUTOMATION_MAX_COUNT
if max_count:
max_count = int(max_count)
if max_count > 0 and Automations.count_by_user(user.id, db=db) >= max_count:
if max_count > 0 and await Automations.count_by_user(user.id, db=db) >= max_count:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=f'Automation limit reached ({max_count})',
@ -90,9 +90,9 @@ def check_automation_limits(request, user, rrule_str: str, db, is_create: bool =
)
def enrich_automation(automation: AutomationModel, db: Session, tz: str = None) -> AutomationResponse:
async def enrich_automation(automation: AutomationModel, db: AsyncSession, tz: str = None) -> AutomationResponse:
"""Full enrichment for single-item views (includes next_runs computation)."""
last_run = AutomationRuns.get_latest(automation.id, db=db)
last_run = await AutomationRuns.get_latest(automation.id, db=db)
return AutomationResponse(
**automation.model_dump(),
last_run=last_run,
@ -112,14 +112,14 @@ async def get_automation_items(
status: Optional[str] = None,
page: Optional[int] = 1,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
check_automations_permission(request, user)
await check_automations_permission(request, user)
limit = PAGE_ITEM_COUNT
page = max(1, page)
skip = (page - 1) * limit
result = Automations.search_automations(
result = await Automations.search_automations(
user_id=user.id,
query=query,
status=status,
@ -130,7 +130,7 @@ async def get_automation_items(
# Batch-fetch latest runs in a single query instead of N+1
ids = [item.id for item in result.items]
latest_runs = AutomationRuns.get_latest_batch(ids, db=db) if ids else {}
latest_runs = await AutomationRuns.get_latest_batch(ids, db=db) if ids else {}
return {
'items': [
@ -154,9 +154,9 @@ async def create_new_automation(
request: Request,
form_data: AutomationForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
check_automations_permission(request, user)
await check_automations_permission(request, user)
try:
validate_rrule(form_data.data.rrule)
except ValueError as e:
@ -165,7 +165,7 @@ async def create_new_automation(
detail=str(e),
)
check_automation_limits(request, user, form_data.data.rrule, db, is_create=True)
await check_automation_limits(request, user, form_data.data.rrule, db, is_create=True)
# Validate terminal server exists if linked
if form_data.data.terminal and form_data.data.terminal.server_id:
@ -177,8 +177,8 @@ async def create_new_automation(
)
tz = user.timezone
automation = Automations.insert(user.id, form_data, next_run_ns(form_data.data.rrule, tz=tz), db=db)
return enrich_automation(automation, db, tz=tz)
automation = await Automations.insert(user.id, form_data, next_run_ns(form_data.data.rrule, tz=tz), db=db)
return await enrich_automation(automation, db, tz=tz)
############################
@ -191,12 +191,12 @@ async def get_automation_by_id(
request: Request,
id: str,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
check_automations_permission(request, user)
automation = Automations.get_by_id(id, db=db)
await check_automations_permission(request, user)
automation = await Automations.get_by_id(id, db=db)
check_automation_access(automation, user)
return enrich_automation(automation, db, tz=user.timezone)
return await enrich_automation(automation, db, tz=user.timezone)
############################
@ -210,10 +210,10 @@ async def update_automation_by_id(
id: str,
form_data: AutomationForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
check_automations_permission(request, user)
automation = Automations.get_by_id(id, db=db)
await check_automations_permission(request, user)
automation = await Automations.get_by_id(id, db=db)
check_automation_access(automation, user)
try:
@ -224,7 +224,7 @@ async def update_automation_by_id(
detail=str(e),
)
check_automation_limits(request, user, form_data.data.rrule, db, is_create=False)
await check_automation_limits(request, user, form_data.data.rrule, db, is_create=False)
# Validate terminal server exists if linked
if form_data.data.terminal and form_data.data.terminal.server_id:
@ -236,8 +236,8 @@ async def update_automation_by_id(
)
tz = user.timezone
updated = Automations.update_by_id(id, form_data, next_run_ns(form_data.data.rrule, tz=tz), db=db)
return enrich_automation(updated, db, tz=tz)
updated = await Automations.update_by_id(id, form_data, next_run_ns(form_data.data.rrule, tz=tz), db=db)
return await enrich_automation(updated, db, tz=tz)
############################
@ -250,13 +250,13 @@ async def toggle_automation_by_id(
request: Request,
id: str,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
check_automations_permission(request, user)
automation = Automations.get_by_id(id, db=db)
await check_automations_permission(request, user)
automation = await Automations.get_by_id(id, db=db)
check_automation_access(automation, user)
toggled = Automations.toggle(id, next_run_ns(automation.data['rrule'], tz=user.timezone), db=db)
return enrich_automation(toggled, db, tz=user.timezone)
toggled = await Automations.toggle(id, next_run_ns(automation.data['rrule'], tz=user.timezone), db=db)
return await enrich_automation(toggled, db, tz=user.timezone)
############################
@ -269,13 +269,13 @@ async def run_automation_by_id(
request: Request,
id: str,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
check_automations_permission(request, user)
automation = Automations.get_by_id(id, db=db)
await check_automations_permission(request, user)
automation = await Automations.get_by_id(id, db=db)
check_automation_access(automation, user)
asyncio.create_task(execute_automation(request.app, automation))
return enrich_automation(automation, db, tz=user.timezone)
return await enrich_automation(automation, db, tz=user.timezone)
############################
@ -288,13 +288,13 @@ async def delete_automation_by_id(
request: Request,
id: str,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
check_automations_permission(request, user)
automation = Automations.get_by_id(id, db=db)
await check_automations_permission(request, user)
automation = await Automations.get_by_id(id, db=db)
check_automation_access(automation, user)
AutomationRuns.delete_by_automation(id, db=db)
return Automations.delete(id, db=db)
await AutomationRuns.delete_by_automation(id, db=db)
return await Automations.delete(id, db=db)
############################
@ -309,9 +309,9 @@ async def get_automation_runs(
skip: int = 0,
limit: int = 50,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
check_automations_permission(request, user)
automation = Automations.get_by_id(id, db=db)
await check_automations_permission(request, user)
automation = await Automations.get_by_id(id, db=db)
check_automation_access(automation, user)
return AutomationRuns.get_by_automation(id, skip=skip, limit=limit, db=db)
return await AutomationRuns.get_by_automation(id, skip=skip, limit=limit, db=db)

File diff suppressed because it is too large Load diff

View file

@ -1,7 +1,7 @@
import json
import logging
from typing import Optional
from sqlalchemy.orm import Session
from sqlalchemy.ext.asyncio import AsyncSession
import asyncio
from fastapi.responses import StreamingResponse
@ -25,7 +25,7 @@ from open_webui.models.chats import (
)
from open_webui.models.tags import TagModel, Tags
from open_webui.models.folders import Folders
from open_webui.internal.db import get_session
from open_webui.internal.db import get_async_session
from open_webui.config import ENABLE_ADMIN_CHAT_ACCESS, ENABLE_ADMIN_EXPORT
from open_webui.constants import ERROR_MESSAGES
@ -49,19 +49,19 @@ router = APIRouter()
@router.get('/', response_model=list[ChatTitleIdResponse])
@router.get('/list', response_model=list[ChatTitleIdResponse])
def get_session_user_chat_list(
async def get_session_user_chat_list(
user=Depends(get_verified_user),
page: Optional[int] = None,
include_pinned: Optional[bool] = False,
include_folders: Optional[bool] = False,
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
try:
if page is not None:
limit = 60
skip = (page - 1) * limit
return Chats.get_chat_title_id_list_by_user_id(
return await Chats.get_chat_title_id_list_by_user_id(
user.id,
include_folders=include_folders,
include_pinned=include_pinned,
@ -70,7 +70,7 @@ def get_session_user_chat_list(
db=db,
)
else:
return Chats.get_chat_title_id_list_by_user_id(
return await Chats.get_chat_title_id_list_by_user_id(
user.id,
include_folders=include_folders,
include_pinned=include_pinned,
@ -88,17 +88,17 @@ def get_session_user_chat_list(
@router.get('/stats/usage', response_model=ChatUsageStatsListResponse)
def get_session_user_chat_usage_stats(
async def get_session_user_chat_usage_stats(
items_per_page: Optional[int] = 50,
page: Optional[int] = 1,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
try:
limit = items_per_page
skip = (page - 1) * limit
result = Chats.get_chats_by_user_id(user.id, skip=skip, limit=limit, db=db)
result = await Chats.get_chats_by_user_id(user.id, skip=skip, limit=limit, db=db)
chats = result.items
total = result.total
@ -332,11 +332,11 @@ def _process_chat_for_export(chat) -> Optional[ChatStatsExport]:
return None
def calculate_chat_stats(user_id, skip=0, limit=10, filter=None):
async def calculate_chat_stats(user_id, skip=0, limit=10, filter=None):
if filter is None:
filter = {}
result = Chats.get_chats_by_user_id(
result = await Chats.get_chats_by_user_id(
user_id,
skip=skip,
limit=limit,
@ -352,12 +352,12 @@ def calculate_chat_stats(user_id, skip=0, limit=10, filter=None):
return chat_stats_export_list, result.total
def generate_chat_stats_jsonl_generator(user_id, filter):
async def generate_chat_stats_jsonl_generator(user_id, filter):
"""
Synchronous generator for streaming chat stats export.
Async generator for streaming chat stats export.
NOTE: We intentionally do NOT pass a shared db session here. Instead, we let
each batch create its own short-lived session via get_db_context(None).
each batch create its own short-lived session via get_async_db_context(None).
This is critical for SQLite in low-resource environments because:
1. SQLite uses file-level locking
2. Holding a session open for the entire streaming duration blocks other requests
@ -368,12 +368,12 @@ def generate_chat_stats_jsonl_generator(user_id, filter):
while True:
# Each batch gets its own session that closes after the query
result = Chats.get_chats_by_user_id(
result = await Chats.get_chats_by_user_id(
user_id,
filter=filter,
skip=skip,
limit=limit,
db=None, # Let get_db_context create a fresh session per batch
db=None, # Let get_async_db_context create a fresh session per batch
)
if not result.items:
break
@ -421,7 +421,7 @@ async def export_chat_stats(
limit = CHAT_EXPORT_PAGE_ITEM_COUNT
skip = (page - 1) * limit
chat_stats_export_list, total = await asyncio.to_thread(calculate_chat_stats, user.id, skip, limit, filter)
chat_stats_export_list, total = await calculate_chat_stats(user.id, skip, limit, filter)
return ChatStatsExportList(items=chat_stats_export_list, total=total, page=page)
@ -440,7 +440,7 @@ async def export_single_chat_stats(
request: Request,
chat_id: str,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""
Export stats for exactly one chat by ID.
@ -454,7 +454,7 @@ async def export_single_chat_stats(
)
try:
chat = Chats.get_chat_by_id(chat_id, db=db)
chat = await Chats.get_chat_by_id(chat_id, db=db)
if not chat:
raise HTTPException(
@ -469,8 +469,8 @@ async def export_single_chat_stats(
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
# Process the chat for export
chat_stats = await asyncio.to_thread(_process_chat_for_export, chat)
# Process the chat for export (pure computation, no DB)
chat_stats = _process_chat_for_export(chat)
if not chat_stats:
raise HTTPException(
@ -491,15 +491,15 @@ async def export_single_chat_stats(
async def delete_all_user_chats(
request: Request,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if user.role == 'user' and not has_permission(user.id, 'chat.delete', request.app.state.config.USER_PERMISSIONS):
if user.role == 'user' and not await has_permission(user.id, 'chat.delete', request.app.state.config.USER_PERMISSIONS):
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
result = Chats.delete_chats_by_user_id(user.id, db=db)
result = await Chats.delete_chats_by_user_id(user.id, db=db)
return result
@ -516,7 +516,7 @@ async def get_user_chat_list_by_user_id(
order_by: Optional[str] = None,
direction: Optional[str] = None,
user=Depends(get_admin_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if not ENABLE_ADMIN_CHAT_ACCESS:
raise HTTPException(
@ -538,7 +538,7 @@ async def get_user_chat_list_by_user_id(
if direction:
filter['direction'] = direction
return Chats.get_chat_list_by_user_id(user_id, include_archived=True, filter=filter, skip=skip, limit=limit, db=db)
return await Chats.get_chat_list_by_user_id(user_id, include_archived=True, filter=filter, skip=skip, limit=limit, db=db)
############################
@ -550,10 +550,10 @@ async def get_user_chat_list_by_user_id(
async def create_new_chat(
form_data: ChatForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
try:
chat = Chats.insert_new_chat(user.id, form_data, db=db)
chat = await Chats.insert_new_chat(user.id, form_data, db=db)
return ChatResponse(**chat.model_dump())
except Exception as e:
log.exception(e)
@ -569,10 +569,10 @@ async def create_new_chat(
async def import_chats(
form_data: ChatsImportForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
try:
chats = Chats.import_chats(user.id, form_data.chats, db=db)
chats = await Chats.import_chats(user.id, form_data.chats, db=db)
return chats
except Exception as e:
log.exception(e)
@ -585,11 +585,11 @@ async def import_chats(
@router.get('/search', response_model=list[ChatTitleIdResponse])
def search_user_chats(
async def search_user_chats(
text: str,
page: Optional[int] = None,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if page is None:
page = 1
@ -599,7 +599,7 @@ def search_user_chats(
chat_list = [
ChatTitleIdResponse(**chat.model_dump())
for chat in Chats.get_chats_by_user_id_and_search_text(user.id, text, skip=skip, limit=limit, db=db)
for chat in await Chats.get_chats_by_user_id_and_search_text(user.id, text, skip=skip, limit=limit, db=db)
]
# Delete tag if no chat is found
@ -607,9 +607,9 @@ def search_user_chats(
if page == 1 and len(words) == 1 and words[0].startswith('tag:'):
tag_id = words[0].replace('tag:', '')
if len(chat_list) == 0:
if Tags.get_tag_by_name_and_user_id(tag_id, user.id, db=db):
if await Tags.get_tag_by_name_and_user_id(tag_id, user.id, db=db):
log.debug(f'deleting tag: {tag_id}')
Tags.delete_tag_by_name_and_user_id(tag_id, user.id, db=db)
await Tags.delete_tag_by_name_and_user_id(tag_id, user.id, db=db)
return chat_list
@ -620,15 +620,15 @@ def search_user_chats(
@router.get('/folder/{folder_id}', response_model=list[ChatResponse])
async def get_chats_by_folder_id(folder_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
async def get_chats_by_folder_id(folder_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
folder_ids = [folder_id]
children_folders = Folders.get_children_folders_by_id_and_user_id(folder_id, user.id, db=db)
children_folders = await Folders.get_children_folders_by_id_and_user_id(folder_id, user.id, db=db)
if children_folders:
folder_ids.extend([folder.id for folder in children_folders])
return [
ChatResponse(**chat.model_dump())
for chat in Chats.get_chats_by_folder_ids_and_user_id(folder_ids, user.id, db=db)
for chat in await Chats.get_chats_by_folder_ids_and_user_id(folder_ids, user.id, db=db)
]
@ -637,13 +637,13 @@ async def get_chat_list_by_folder_id(
folder_id: str,
page: Optional[int] = 1,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
try:
limit = 10
skip = (page - 1) * limit
chats = Chats.get_chats_by_folder_id_and_user_id(folder_id, user.id, skip=skip, limit=limit, db=db)
chats = await Chats.get_chats_by_folder_id_and_user_id(folder_id, user.id, skip=skip, limit=limit, db=db)
return [
{'title': chat.title, 'id': chat.id, 'updated_at': chat.updated_at, 'last_read_at': chat.last_read_at}
for chat in chats
@ -660,8 +660,8 @@ async def get_chat_list_by_folder_id(
@router.get('/pinned', response_model=list[ChatTitleIdResponse])
async def get_user_pinned_chats(user=Depends(get_verified_user), db: Session = Depends(get_session)):
return Chats.get_pinned_chats_by_user_id(user.id, db=db)
async def get_user_pinned_chats(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
return await Chats.get_pinned_chats_by_user_id(user.id, db=db)
############################
@ -670,8 +670,8 @@ async def get_user_pinned_chats(user=Depends(get_verified_user), db: Session = D
@router.get('/all', response_model=list[ChatResponse])
async def get_user_chats(user=Depends(get_verified_user), db: Session = Depends(get_session)):
result = Chats.get_chats_by_user_id(user.id, db=db)
async def get_user_chats(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
result = await Chats.get_chats_by_user_id(user.id, db=db)
return [ChatResponse(**chat.model_dump()) for chat in result.items]
@ -681,8 +681,8 @@ async def get_user_chats(user=Depends(get_verified_user), db: Session = Depends(
@router.get('/all/archived', response_model=list[ChatResponse])
async def get_user_archived_chats(user=Depends(get_verified_user), db: Session = Depends(get_session)):
return [ChatResponse(**chat.model_dump()) for chat in Chats.get_archived_chats_by_user_id(user.id, db=db)]
async def get_user_archived_chats(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
return [ChatResponse(**chat.model_dump()) for chat in await Chats.get_archived_chats_by_user_id(user.id, db=db)]
############################
@ -691,9 +691,9 @@ async def get_user_archived_chats(user=Depends(get_verified_user), db: Session =
@router.get('/all/tags', response_model=list[TagModel])
async def get_all_user_tags(user=Depends(get_verified_user), db: Session = Depends(get_session)):
async def get_all_user_tags(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
try:
tags = Tags.get_tags_by_user_id(user.id, db=db)
tags = await Tags.get_tags_by_user_id(user.id, db=db)
return tags
except Exception as e:
log.exception(e)
@ -706,13 +706,13 @@ async def get_all_user_tags(user=Depends(get_verified_user), db: Session = Depen
@router.get('/all/db', response_model=list[ChatResponse])
async def get_all_user_chats_in_db(user=Depends(get_admin_user), db: Session = Depends(get_session)):
async def get_all_user_chats_in_db(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
if not ENABLE_ADMIN_EXPORT:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
return [ChatResponse(**chat.model_dump()) for chat in Chats.get_chats(db=db)]
return [ChatResponse(**chat.model_dump()) for chat in await Chats.get_chats(db=db)]
############################
@ -727,7 +727,7 @@ async def get_archived_session_user_chat_list(
order_by: Optional[str] = None,
direction: Optional[str] = None,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if page is None:
page = 1
@ -743,7 +743,7 @@ async def get_archived_session_user_chat_list(
if direction:
filter['direction'] = direction
return Chats.get_archived_chat_list_by_user_id(
return await Chats.get_archived_chat_list_by_user_id(
user.id,
filter=filter,
skip=skip,
@ -758,8 +758,8 @@ async def get_archived_session_user_chat_list(
@router.post('/archive/all', response_model=bool)
async def archive_all_chats(user=Depends(get_verified_user), db: Session = Depends(get_session)):
return Chats.archive_all_chats_by_user_id(user.id, db=db)
async def archive_all_chats(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
return await Chats.archive_all_chats_by_user_id(user.id, db=db)
############################
@ -768,8 +768,8 @@ async def archive_all_chats(user=Depends(get_verified_user), db: Session = Depen
@router.post('/unarchive/all', response_model=bool)
async def unarchive_all_chats(user=Depends(get_verified_user), db: Session = Depends(get_session)):
return Chats.unarchive_all_chats_by_user_id(user.id, db=db)
async def unarchive_all_chats(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
return await Chats.unarchive_all_chats_by_user_id(user.id, db=db)
############################
@ -784,7 +784,7 @@ async def get_shared_session_user_chat_list(
order_by: Optional[str] = None,
direction: Optional[str] = None,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if page is None:
page = 1
@ -800,7 +800,7 @@ async def get_shared_session_user_chat_list(
if direction:
filter['direction'] = direction
return Chats.get_shared_chat_list_by_user_id(
return await Chats.get_shared_chat_list_by_user_id(
user.id,
filter=filter,
skip=skip,
@ -815,14 +815,14 @@ async def get_shared_session_user_chat_list(
@router.get('/share/{share_id}', response_model=Optional[ChatResponse])
async def get_shared_chat_by_id(share_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
async def get_shared_chat_by_id(share_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
if user.role == 'pending':
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.NOT_FOUND)
if user.role == 'user' or (user.role == 'admin' and not ENABLE_ADMIN_CHAT_ACCESS):
chat = Chats.get_chat_by_share_id(share_id, db=db)
chat = await Chats.get_chat_by_share_id(share_id, db=db)
elif user.role == 'admin' and ENABLE_ADMIN_CHAT_ACCESS:
chat = Chats.get_chat_by_id(share_id, db=db)
chat = await Chats.get_chat_by_id(share_id, db=db)
if chat:
return ChatResponse(**chat.model_dump())
@ -849,11 +849,11 @@ class TagFilterForm(TagForm):
async def get_user_chat_list_by_tag_name(
form_data: TagFilterForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
chats = Chats.get_chat_list_by_user_id_and_tag_name(user.id, form_data.name, form_data.skip, form_data.limit, db=db)
chats = await Chats.get_chat_list_by_user_id_and_tag_name(user.id, form_data.name, form_data.skip, form_data.limit, db=db)
if len(chats) == 0:
Tags.delete_tag_by_name_and_user_id(form_data.name, user.id, db=db)
await Tags.delete_tag_by_name_and_user_id(form_data.name, user.id, db=db)
return chats
@ -864,8 +864,8 @@ async def get_user_chat_list_by_tag_name(
@router.get('/{id}', response_model=Optional[ChatResponse])
async def get_chat_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
async def get_chat_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
if chat:
return ChatResponse(**chat.model_dump())
@ -884,12 +884,12 @@ async def update_chat_by_id(
id: str,
form_data: ChatForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
if chat:
updated_chat = {**chat.chat, **form_data.chat}
chat = Chats.update_chat_by_id(id, updated_chat, db=db)
chat = await Chats.update_chat_by_id(id, updated_chat, db=db)
return ChatResponse(**chat.model_dump())
else:
raise HTTPException(
@ -911,9 +911,9 @@ async def update_chat_message_by_id(
message_id: str,
form_data: MessageForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
chat = Chats.get_chat_by_id(id, db=db)
chat = await Chats.get_chat_by_id(id, db=db)
if not chat:
raise HTTPException(
@ -927,7 +927,7 @@ async def update_chat_message_by_id(
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
chat = Chats.upsert_message_to_chat_by_id_and_message_id(
chat = await Chats.upsert_message_to_chat_by_id_and_message_id(
id,
message_id,
{
@ -935,7 +935,7 @@ async def update_chat_message_by_id(
},
)
event_emitter = get_event_emitter(
event_emitter = await get_event_emitter(
{
'user_id': user.id,
'chat_id': id,
@ -973,9 +973,9 @@ async def send_chat_message_event_by_id(
message_id: str,
form_data: EventForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
chat = Chats.get_chat_by_id(id, db=db)
chat = await Chats.get_chat_by_id(id, db=db)
if not chat:
raise HTTPException(
@ -989,7 +989,7 @@ async def send_chat_message_event_by_id(
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
event_emitter = get_event_emitter(
event_emitter = await get_event_emitter(
{
'user_id': user.id,
'chat_id': id,
@ -1017,36 +1017,36 @@ async def delete_chat_by_id(
request: Request,
id: str,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if user.role == 'admin':
chat = Chats.get_chat_by_id(id, db=db)
chat = await Chats.get_chat_by_id(id, db=db)
if not chat:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.NOT_FOUND,
)
Chats.delete_orphan_tags_for_user(chat.meta.get('tags', []), user.id, threshold=1, db=db)
await Chats.delete_orphan_tags_for_user(chat.meta.get('tags', []), user.id, threshold=1, db=db)
result = Chats.delete_chat_by_id(id, db=db)
result = await Chats.delete_chat_by_id(id, db=db)
return result
else:
if not has_permission(user.id, 'chat.delete', request.app.state.config.USER_PERMISSIONS):
if not await has_permission(user.id, 'chat.delete', request.app.state.config.USER_PERMISSIONS):
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
if not chat:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.NOT_FOUND,
)
Chats.delete_orphan_tags_for_user(chat.meta.get('tags', []), user.id, threshold=1, db=db)
await Chats.delete_orphan_tags_for_user(chat.meta.get('tags', []), user.id, threshold=1, db=db)
result = Chats.delete_chat_by_id_and_user_id(id, user.id, db=db)
result = await Chats.delete_chat_by_id_and_user_id(id, user.id, db=db)
return result
@ -1056,8 +1056,8 @@ async def delete_chat_by_id(
@router.get('/{id}/pinned', response_model=Optional[bool])
async def get_pinned_status_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
async def get_pinned_status_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
if chat:
return chat.pinned
else:
@ -1070,10 +1070,10 @@ async def get_pinned_status_by_id(id: str, user=Depends(get_verified_user), db:
@router.post('/{id}/pin', response_model=Optional[ChatResponse])
async def pin_chat_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
async def pin_chat_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
if chat:
chat = Chats.toggle_chat_pinned_by_id(id, db=db)
chat = await Chats.toggle_chat_pinned_by_id(id, db=db)
return chat
else:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.DEFAULT())
@ -1093,9 +1093,9 @@ async def clone_chat_by_id(
form_data: CloneForm,
id: str,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
if chat:
updated_chat = {
**chat.chat,
@ -1104,7 +1104,7 @@ async def clone_chat_by_id(
'title': form_data.title if form_data.title else f'Clone of {chat.title}',
}
chats = Chats.import_chats(
chats = await Chats.import_chats(
user.id,
[
ChatImportForm(
@ -1137,11 +1137,11 @@ async def clone_chat_by_id(
@router.post('/{id}/clone/shared', response_model=Optional[ChatResponse])
async def clone_shared_chat_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
async def clone_shared_chat_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
if user.role == 'admin':
chat = Chats.get_chat_by_id(id, db=db)
chat = await Chats.get_chat_by_id(id, db=db)
else:
chat = Chats.get_chat_by_share_id(id, db=db)
chat = await Chats.get_chat_by_share_id(id, db=db)
if chat:
updated_chat = {
@ -1151,7 +1151,7 @@ async def clone_shared_chat_by_id(id: str, user=Depends(get_verified_user), db:
'title': f'Clone of {chat.title}',
}
chats = Chats.import_chats(
chats = await Chats.import_chats(
user.id,
[
ChatImportForm(
@ -1184,18 +1184,18 @@ async def clone_shared_chat_by_id(id: str, user=Depends(get_verified_user), db:
@router.post('/{id}/archive', response_model=Optional[ChatResponse])
async def archive_chat_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
async def archive_chat_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
if chat:
chat = Chats.toggle_chat_archive_by_id(id, db=db)
chat = await Chats.toggle_chat_archive_by_id(id, db=db)
tag_ids = chat.meta.get('tags', [])
if chat.archived:
# Archived chats are excluded from count — clean up orphans
Chats.delete_orphan_tags_for_user(tag_ids, user.id, db=db)
await Chats.delete_orphan_tags_for_user(tag_ids, user.id, db=db)
else:
# Unarchived — ensure tag rows exist
Tags.ensure_tags_exist(tag_ids, user.id, db=db)
await Tags.ensure_tags_exist(tag_ids, user.id, db=db)
return ChatResponse(**chat.model_dump())
else:
@ -1212,24 +1212,24 @@ async def share_chat_by_id(
request: Request,
id: str,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if (user.role != 'admin') and (
not has_permission(user.id, 'chat.share', request.app.state.config.USER_PERMISSIONS)
not await has_permission(user.id, 'chat.share', request.app.state.config.USER_PERMISSIONS)
):
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
if chat:
if chat.share_id:
shared_chat = Chats.update_shared_chat_by_chat_id(chat.id, db=db)
shared_chat = await Chats.update_shared_chat_by_chat_id(chat.id, db=db)
return ChatResponse(**shared_chat.model_dump())
shared_chat = Chats.insert_shared_chat_by_chat_id(chat.id, db=db)
shared_chat = await Chats.insert_shared_chat_by_chat_id(chat.id, db=db)
if not shared_chat:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
@ -1250,14 +1250,14 @@ async def share_chat_by_id(
@router.delete('/{id}/share', response_model=Optional[bool])
async def delete_shared_chat_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
async def delete_shared_chat_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
if chat:
if not chat.share_id:
return False
result = Chats.delete_shared_chat_by_chat_id(id, db=db)
update_result = Chats.update_chat_share_id_by_id(id, None, db=db)
result = await Chats.delete_shared_chat_by_chat_id(id, db=db)
update_result = await Chats.update_chat_share_id_by_id(id, None, db=db)
return result and update_result != None
else:
@ -1281,11 +1281,11 @@ async def update_chat_folder_id_by_id(
id: str,
form_data: ChatFolderIdForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
if chat:
chat = Chats.update_chat_folder_id_by_id_and_user_id(id, user.id, form_data.folder_id, db=db)
chat = await Chats.update_chat_folder_id_by_id_and_user_id(id, user.id, form_data.folder_id, db=db)
return ChatResponse(**chat.model_dump())
else:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.DEFAULT())
@ -1297,11 +1297,11 @@ async def update_chat_folder_id_by_id(
@router.get('/{id}/tags', response_model=list[TagModel])
async def get_chat_tags_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
async def get_chat_tags_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
if chat:
tags = chat.meta.get('tags', [])
return Tags.get_tags_by_ids_and_user_id(tags, user.id, db=db)
return await Tags.get_tags_by_ids_and_user_id(tags, user.id, db=db)
else:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.NOT_FOUND)
@ -1316,9 +1316,9 @@ async def add_tag_by_id_and_tag_name(
id: str,
form_data: TagForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
if chat:
tags = chat.meta.get('tags', [])
tag_id = form_data.name.replace(' ', '_').lower()
@ -1330,11 +1330,11 @@ async def add_tag_by_id_and_tag_name(
)
if tag_id not in tags:
Chats.add_chat_tag_by_id_and_user_id_and_tag_name(id, user.id, form_data.name, db=db)
await Chats.add_chat_tag_by_id_and_user_id_and_tag_name(id, user.id, form_data.name, db=db)
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
tags = chat.meta.get('tags', [])
return Tags.get_tags_by_ids_and_user_id(tags, user.id, db=db)
return await Tags.get_tags_by_ids_and_user_id(tags, user.id, db=db)
else:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.DEFAULT())
@ -1349,18 +1349,18 @@ async def delete_tag_by_id_and_tag_name(
id: str,
form_data: TagForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
if chat:
Chats.delete_tag_by_id_and_user_id_and_tag_name(id, user.id, form_data.name, db=db)
await Chats.delete_tag_by_id_and_user_id_and_tag_name(id, user.id, form_data.name, db=db)
if Chats.count_chats_by_tag_name_and_user_id(form_data.name, user.id, db=db) == 0:
Tags.delete_tag_by_name_and_user_id(form_data.name, user.id, db=db)
if await Chats.count_chats_by_tag_name_and_user_id(form_data.name, user.id, db=db) == 0:
await Tags.delete_tag_by_name_and_user_id(form_data.name, user.id, db=db)
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
tags = chat.meta.get('tags', [])
return Tags.get_tags_by_ids_and_user_id(tags, user.id, db=db)
return await Tags.get_tags_by_ids_and_user_id(tags, user.id, db=db)
else:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.NOT_FOUND)
@ -1371,12 +1371,12 @@ async def delete_tag_by_id_and_tag_name(
@router.delete('/{id}/tags/all', response_model=Optional[bool])
async def delete_all_tags_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
chat = Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
async def delete_all_tags_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
if chat:
old_tags = chat.meta.get('tags', [])
Chats.delete_all_tags_by_id_and_user_id(id, user.id, db=db)
Chats.delete_orphan_tags_for_user(old_tags, user.id, db=db)
await Chats.delete_all_tags_by_id_and_user_id(id, user.id, db=db)
await Chats.delete_orphan_tags_for_user(old_tags, user.id, db=db)
return True
else:

View file

@ -20,8 +20,8 @@ from open_webui.models.feedbacks import (
from open_webui.constants import ERROR_MESSAGES
from open_webui.utils.auth import get_admin_user, get_verified_user
from open_webui.internal.db import get_session
from sqlalchemy.orm import Session
from open_webui.internal.db import get_async_session
from sqlalchemy.ext.asyncio import AsyncSession
log = logging.getLogger(__name__)
@ -208,10 +208,10 @@ class LeaderboardResponse(BaseModel):
async def get_leaderboard(
query: Optional[str] = None,
user=Depends(get_admin_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""Get model leaderboard with Elo ratings. Query filters by tag similarity."""
feedbacks = Feedbacks.get_feedbacks_for_leaderboard(db=db)
feedbacks = await Feedbacks.get_feedbacks_for_leaderboard(db=db)
similarities = None
if query and query.strip():
@ -244,10 +244,10 @@ async def get_model_history(
model_id: str,
days: int = 30,
user=Depends(get_admin_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""Get daily win/loss history for a specific model."""
history = Feedbacks.get_model_evaluation_history(model_id=model_id, days=days, db=db)
history = await Feedbacks.get_model_evaluation_history(model_id=model_id, days=days, db=db)
return ModelHistoryResponse(model_id=model_id, history=history)
@ -292,24 +292,24 @@ async def update_config(
@router.get('/feedbacks/models', response_model=list[str])
async def get_feedback_model_ids(user=Depends(get_admin_user), db: Session = Depends(get_session)):
return Feedbacks.get_distinct_model_ids(db=db)
async def get_feedback_model_ids(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
return await Feedbacks.get_distinct_model_ids(db=db)
@router.get('/feedbacks/all', response_model=list[FeedbackResponse])
async def get_all_feedbacks(user=Depends(get_admin_user), db: Session = Depends(get_session)):
feedbacks = Feedbacks.get_all_feedbacks(db=db)
async def get_all_feedbacks(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
feedbacks = await Feedbacks.get_all_feedbacks(db=db)
return feedbacks
@router.get('/feedbacks/all/ids', response_model=list[FeedbackIdResponse])
async def get_all_feedback_ids(user=Depends(get_admin_user), db: Session = Depends(get_session)):
return Feedbacks.get_all_feedback_ids(db=db)
async def get_all_feedback_ids(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
return await Feedbacks.get_all_feedback_ids(db=db)
@router.delete('/feedbacks/all')
async def delete_all_feedbacks(user=Depends(get_admin_user), db: Session = Depends(get_session)):
success = Feedbacks.delete_all_feedbacks(db=db)
async def delete_all_feedbacks(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
success = await Feedbacks.delete_all_feedbacks(db=db)
return success
@ -317,23 +317,23 @@ async def delete_all_feedbacks(user=Depends(get_admin_user), db: Session = Depen
async def export_all_feedbacks(
model_id: Optional[str] = None,
user=Depends(get_admin_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
feedbacks = Feedbacks.get_all_feedbacks(db=db)
feedbacks = await Feedbacks.get_all_feedbacks(db=db)
if model_id:
feedbacks = [f for f in feedbacks if f.data and f.data.get('model_id') == model_id]
return feedbacks
@router.get('/feedbacks/user', response_model=list[FeedbackUserResponse])
async def get_feedbacks(user=Depends(get_verified_user), db: Session = Depends(get_session)):
feedbacks = Feedbacks.get_feedbacks_by_user_id(user.id, db=db)
async def get_feedbacks(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
feedbacks = await Feedbacks.get_feedbacks_by_user_id(user.id, db=db)
return feedbacks
@router.delete('/feedbacks', response_model=bool)
async def delete_feedbacks(user=Depends(get_verified_user), db: Session = Depends(get_session)):
success = Feedbacks.delete_feedbacks_by_user_id(user.id, db=db)
async def delete_feedbacks(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
success = await Feedbacks.delete_feedbacks_by_user_id(user.id, db=db)
return success
@ -347,7 +347,7 @@ async def get_feedbacks(
page: Optional[int] = 1,
model_id: Optional[str] = None,
user=Depends(get_admin_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
limit = PAGE_ITEM_COUNT
@ -362,7 +362,7 @@ async def get_feedbacks(
if model_id:
filter['model_id'] = model_id
result = Feedbacks.get_feedback_items(filter=filter, skip=skip, limit=limit, db=db)
result = await Feedbacks.get_feedback_items(filter=filter, skip=skip, limit=limit, db=db)
return result
@ -371,9 +371,9 @@ async def create_feedback(
request: Request,
form_data: FeedbackForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
feedback = Feedbacks.insert_new_feedback(user_id=user.id, form_data=form_data, db=db)
feedback = await Feedbacks.insert_new_feedback(user_id=user.id, form_data=form_data, db=db)
if not feedback:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
@ -384,11 +384,11 @@ async def create_feedback(
@router.get('/feedback/{id}', response_model=FeedbackModel)
async def get_feedback_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
async def get_feedback_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
if user.role == 'admin':
feedback = Feedbacks.get_feedback_by_id(id=id, db=db)
feedback = await Feedbacks.get_feedback_by_id(id=id, db=db)
else:
feedback = Feedbacks.get_feedback_by_id_and_user_id(id=id, user_id=user.id, db=db)
feedback = await Feedbacks.get_feedback_by_id_and_user_id(id=id, user_id=user.id, db=db)
if not feedback:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND)
@ -401,12 +401,12 @@ async def update_feedback_by_id(
id: str,
form_data: FeedbackForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if user.role == 'admin':
feedback = Feedbacks.update_feedback_by_id(id=id, form_data=form_data, db=db)
feedback = await Feedbacks.update_feedback_by_id(id=id, form_data=form_data, db=db)
else:
feedback = Feedbacks.update_feedback_by_id_and_user_id(id=id, user_id=user.id, form_data=form_data, db=db)
feedback = await Feedbacks.update_feedback_by_id_and_user_id(id=id, user_id=user.id, form_data=form_data, db=db)
if not feedback:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND)
@ -415,11 +415,11 @@ async def update_feedback_by_id(
@router.delete('/feedback/{id}')
async def delete_feedback_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
async def delete_feedback_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
if user.role == 'admin':
success = Feedbacks.delete_feedback_by_id(id=id, db=db)
success = await Feedbacks.delete_feedback_by_id(id=id, db=db)
else:
success = Feedbacks.delete_feedback_by_id_and_user_id(id=id, user_id=user.id, db=db)
success = await Feedbacks.delete_feedback_by_id_and_user_id(id=id, user_id=user.id, db=db)
if not success:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND)

View file

@ -21,8 +21,8 @@ from fastapi import (
)
from fastapi.responses import FileResponse, StreamingResponse
from sqlalchemy.orm import Session
from open_webui.internal.db import get_session, SessionLocal
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import get_async_session, SessionLocal
from open_webui.constants import ERROR_MESSAGES
from open_webui.retrieval.vector.factory import VECTOR_DB_CLIENT
@ -88,16 +88,16 @@ def _is_text_file(file_path: str, chunk_size: int = 8192) -> bool:
return False
def process_uploaded_file(
async def process_uploaded_file(
request,
file,
file_path,
file_item,
file_metadata,
user,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
):
def _process_handler(db_session):
async def _process_handler(db_session):
try:
content_type = file.content_type
@ -141,7 +141,7 @@ def process_uploaded_file(
except Exception as e:
log.error(f'Error processing file: {file_item.id}')
Files.update_file_data_by_id(
await Files.update_file_data_by_id(
file_item.id,
{
'status': 'failed',
@ -158,7 +158,7 @@ def process_uploaded_file(
@router.post('/', response_model=FileModelResponse)
def upload_file(
async def upload_file(
request: Request,
background_tasks: BackgroundTasks,
file: UploadFile = File(...),
@ -166,9 +166,9 @@ def upload_file(
process: bool = Query(True),
process_in_background: bool = Query(True),
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
return upload_file_handler(
return await upload_file_handler(
request,
file=file,
metadata=metadata,
@ -180,7 +180,7 @@ def upload_file(
)
def upload_file_handler(
async def upload_file_handler(
request: Request,
file: UploadFile = File(...),
metadata: Optional[dict | str] = Form(None),
@ -188,7 +188,7 @@ def upload_file_handler(
process_in_background: bool = Query(True),
user=Depends(get_verified_user),
background_tasks: Optional[BackgroundTasks] = None,
db: Optional[Session] = None,
db: Optional[AsyncSession] = None,
):
log.info(f'file.content_type: {file.content_type} {process}')
@ -236,7 +236,7 @@ def upload_file_handler(
},
)
file_item = Files.insert_new_file(
file_item = await Files.insert_new_file(
user.id,
FileForm(
**{
@ -258,9 +258,9 @@ def upload_file_handler(
)
if 'channel_id' in file_metadata:
channel = Channels.get_channel_by_id_and_user_id(file_metadata['channel_id'], user.id, db=db)
channel = await Channels.get_channel_by_id_and_user_id(file_metadata['channel_id'], user.id, db=db)
if channel:
Channels.add_file_to_channel_by_id(channel.id, file_item.id, user.id, db=db)
await Channels.add_file_to_channel_by_id(channel.id, file_item.id, user.id, db=db)
if process:
if background_tasks and process_in_background:
@ -275,7 +275,7 @@ def upload_file_handler(
)
return {'status': True, **file_item.model_dump()}
else:
process_uploaded_file(
await process_uploaded_file(
request,
file,
file_path,
@ -317,12 +317,12 @@ async def list_files(
user=Depends(get_verified_user),
page: int = Query(1, ge=1, description='Page number (1-indexed)'),
content: bool = Query(True),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
skip = (page - 1) * PAGE_SIZE
user_id = None if (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL) else user.id
result = Files.get_file_list(user_id=user_id, skip=skip, limit=PAGE_SIZE, db=db)
result = await Files.get_file_list(user_id=user_id, skip=skip, limit=PAGE_SIZE, db=db)
if not content:
for file in result.items:
@ -347,7 +347,7 @@ async def search_files(
skip: int = Query(0, ge=0, description='Number of files to skip'),
limit: int = Query(100, ge=1, le=1000, description='Maximum number of files to return'),
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""
Search for files by filename with support for wildcard patterns.
@ -357,7 +357,7 @@ async def search_files(
user_id = None if (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL) else user.id
# Use optimized database query with pagination
files = Files.search_files(
files = await Files.search_files(
user_id=user_id,
filename=filename,
skip=skip,
@ -385,8 +385,8 @@ async def search_files(
@router.delete('/all')
async def delete_all_files(user=Depends(get_admin_user), db: Session = Depends(get_session)):
result = Files.delete_all_files(db=db)
async def delete_all_files(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
result = await Files.delete_all_files(db=db)
if result:
try:
Storage.delete_all_files()
@ -412,8 +412,8 @@ async def delete_all_files(user=Depends(get_admin_user), db: Session = Depends(g
@router.get('/{id}', response_model=Optional[FileModel])
async def get_file_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
file = Files.get_file_by_id(id, db=db)
async def get_file_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
file = await Files.get_file_by_id(id, db=db)
if not file:
raise HTTPException(
@ -421,7 +421,7 @@ async def get_file_by_id(id: str, user=Depends(get_verified_user), db: Session =
detail=ERROR_MESSAGES.NOT_FOUND,
)
if file.user_id == user.id or user.role == 'admin' or has_access_to_file(id, 'read', user, db=db):
if file.user_id == user.id or user.role == 'admin' or await has_access_to_file(id, 'read', user, db=db):
return file
else:
raise HTTPException(
@ -435,9 +435,9 @@ async def get_file_process_status(
id: str,
stream: bool = Query(False),
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
file = Files.get_file_by_id(id, db=db)
file = await Files.get_file_by_id(id, db=db)
if not file:
raise HTTPException(
@ -445,7 +445,7 @@ async def get_file_process_status(
detail=ERROR_MESSAGES.NOT_FOUND,
)
if file.user_id == user.id or user.role == 'admin' or has_access_to_file(id, 'read', user, db=db):
if file.user_id == user.id or user.role == 'admin' or await has_access_to_file(id, 'read', user, db=db):
if stream:
MAX_FILE_PROCESSING_DURATION = 3600 * 2
@ -454,7 +454,7 @@ async def get_file_process_status(
# Each poll creates its own short-lived session to avoid holding a
# connection for hours. A WebSocket push would be more efficient.
for _ in range(MAX_FILE_PROCESSING_DURATION):
file_item = Files.get_file_by_id(file_id) # Creates own session
file_item = await Files.get_file_by_id(file_id) # Creates own session
if file_item:
data = file_item.model_dump().get('data', {})
status = data.get('status')
@ -495,8 +495,8 @@ async def get_file_process_status(
@router.get('/{id}/data/content')
async def get_file_data_content_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
file = Files.get_file_by_id(id, db=db)
async def get_file_data_content_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
file = await Files.get_file_by_id(id, db=db)
if not file:
raise HTTPException(
@ -504,7 +504,7 @@ async def get_file_data_content_by_id(id: str, user=Depends(get_verified_user),
detail=ERROR_MESSAGES.NOT_FOUND,
)
if file.user_id == user.id or user.role == 'admin' or has_access_to_file(id, 'read', user, db=db):
if file.user_id == user.id or user.role == 'admin' or await has_access_to_file(id, 'read', user, db=db):
return {'content': file.data.get('content', '')}
else:
raise HTTPException(
@ -523,14 +523,14 @@ class ContentForm(BaseModel):
@router.post('/{id}/data/content/update')
def update_file_data_content_by_id(
async def update_file_data_content_by_id(
request: Request,
id: str,
form_data: ContentForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
file = Files.get_file_by_id(id, db=db)
file = await Files.get_file_by_id(id, db=db)
if not file:
raise HTTPException(
@ -538,7 +538,7 @@ def update_file_data_content_by_id(
detail=ERROR_MESSAGES.NOT_FOUND,
)
if file.user_id == user.id or user.role == 'admin' or has_access_to_file(id, 'write', user, db=db):
if file.user_id == user.id or user.role == 'admin' or await has_access_to_file(id, 'write', user, db=db):
try:
process_file(
request,
@ -546,7 +546,7 @@ def update_file_data_content_by_id(
user=user,
db=db,
)
file = Files.get_file_by_id(id=id, db=db)
file = await Files.get_file_by_id(id=id, db=db)
except Exception as e:
log.exception(e)
log.error(f'Error processing file: {file.id}')
@ -554,7 +554,7 @@ def update_file_data_content_by_id(
# Propagate content change to all knowledge collections referencing
# this file. Without this the old embeddings remain in the knowledge
# collection and RAG returns both stale and current data (#20558).
knowledges = Knowledges.get_knowledges_by_file_id(id, db=db)
knowledges = await Knowledges.get_knowledges_by_file_id(id, db=db)
for knowledge in knowledges:
try:
# Remove old embeddings for this file from the KB collection
@ -587,9 +587,9 @@ async def get_file_content_by_id(
id: str,
user=Depends(get_verified_user),
attachment: bool = Query(False),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
file = Files.get_file_by_id(id, db=db)
file = await Files.get_file_by_id(id, db=db)
if not file:
raise HTTPException(
@ -597,7 +597,7 @@ async def get_file_content_by_id(
detail=ERROR_MESSAGES.NOT_FOUND,
)
if file.user_id == user.id or user.role == 'admin' or has_access_to_file(id, 'read', user, db=db):
if file.user_id == user.id or user.role == 'admin' or await has_access_to_file(id, 'read', user, db=db):
try:
file_path = Storage.get_file(file.path)
file_path = Path(file_path)
@ -646,8 +646,8 @@ async def get_file_content_by_id(
@router.get('/{id}/content/html')
async def get_html_file_content_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
file = Files.get_file_by_id(id, db=db)
async def get_html_file_content_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
file = await Files.get_file_by_id(id, db=db)
if not file:
raise HTTPException(
@ -655,14 +655,14 @@ async def get_html_file_content_by_id(id: str, user=Depends(get_verified_user),
detail=ERROR_MESSAGES.NOT_FOUND,
)
file_user = Users.get_user_by_id(file.user_id, db=db)
file_user = await Users.get_user_by_id(file.user_id, db=db)
if not file_user or file_user.role != 'admin':
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.NOT_FOUND,
)
if file.user_id == user.id or user.role == 'admin' or has_access_to_file(id, 'read', user, db=db):
if file.user_id == user.id or user.role == 'admin' or await has_access_to_file(id, 'read', user, db=db):
try:
file_path = Storage.get_file(file.path)
file_path = Path(file_path)
@ -693,8 +693,8 @@ async def get_html_file_content_by_id(id: str, user=Depends(get_verified_user),
@router.get('/{id}/content/{file_name}')
async def get_file_content_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
file = Files.get_file_by_id(id, db=db)
async def get_file_content_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
file = await Files.get_file_by_id(id, db=db)
if not file:
raise HTTPException(
@ -702,7 +702,7 @@ async def get_file_content_by_id(id: str, user=Depends(get_verified_user), db: S
detail=ERROR_MESSAGES.NOT_FOUND,
)
if file.user_id == user.id or user.role == 'admin' or has_access_to_file(id, 'read', user, db=db):
if file.user_id == user.id or user.role == 'admin' or await has_access_to_file(id, 'read', user, db=db):
file_path = file.path
# Handle Unicode filenames
@ -749,8 +749,8 @@ async def get_file_content_by_id(id: str, user=Depends(get_verified_user), db: S
@router.delete('/{id}')
async def delete_file_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
file = Files.get_file_by_id(id, db=db)
async def delete_file_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
file = await Files.get_file_by_id(id, db=db)
if not file:
raise HTTPException(
@ -758,12 +758,12 @@ async def delete_file_by_id(id: str, user=Depends(get_verified_user), db: Sessio
detail=ERROR_MESSAGES.NOT_FOUND,
)
if file.user_id == user.id or user.role == 'admin' or has_access_to_file(id, 'write', user, db=db):
if file.user_id == user.id or user.role == 'admin' or await has_access_to_file(id, 'write', user, db=db):
# Clean up KB associations and embeddings before deleting
knowledges = Knowledges.get_knowledges_by_file_id(id, db=db)
knowledges = await Knowledges.get_knowledges_by_file_id(id, db=db)
for knowledge in knowledges:
# Remove KB-file relationship
Knowledges.remove_file_from_knowledge_by_id(knowledge.id, id, db=db)
await Knowledges.remove_file_from_knowledge_by_id(knowledge.id, id, db=db)
# Clean KB embeddings (same logic as /knowledge/{id}/file/remove)
try:
VECTOR_DB_CLIENT.delete(collection_name=knowledge.id, filter={'file_id': id})
@ -772,7 +772,7 @@ async def delete_file_by_id(id: str, user=Depends(get_verified_user), db: Sessio
except Exception as e:
log.debug(f'KB embedding cleanup for {knowledge.id}: {e}')
result = Files.delete_file_by_id(id, db=db)
result = await Files.delete_file_by_id(id, db=db)
if result:
try:
Storage.delete_file(file.path)

View file

@ -22,8 +22,8 @@ from open_webui.models.knowledge import Knowledges
from open_webui.config import UPLOAD_DIR
from open_webui.constants import ERROR_MESSAGES
from open_webui.internal.db import get_session
from sqlalchemy.orm import Session
from open_webui.internal.db import get_async_session
from sqlalchemy.ext.asyncio import AsyncSession
from fastapi import APIRouter, Depends, File, HTTPException, UploadFile, status, Request
@ -48,7 +48,7 @@ router = APIRouter()
async def get_folders(
request: Request,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if request.app.state.config.ENABLE_FOLDERS is False:
raise HTTPException(
@ -56,7 +56,7 @@ async def get_folders(
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
if user.role != 'admin' and not has_permission(
if user.role != 'admin' and not await has_permission(
user.id,
'features.folders',
request.app.state.config.USER_PERMISSIONS,
@ -67,29 +67,29 @@ async def get_folders(
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
folders = Folders.get_folders_by_user_id(user.id, db=db)
folders = await Folders.get_folders_by_user_id(user.id, db=db)
# Verify folder data integrity
folder_list = []
for folder in folders:
if folder.parent_id and not Folders.get_folder_by_id_and_user_id(folder.parent_id, user.id, db=db):
folder = Folders.update_folder_parent_id_by_id_and_user_id(folder.id, user.id, None, db=db)
if folder.parent_id and not await Folders.get_folder_by_id_and_user_id(folder.parent_id, user.id, db=db):
folder = await Folders.update_folder_parent_id_by_id_and_user_id(folder.id, user.id, None, db=db)
if folder.data:
if 'files' in folder.data:
valid_files = []
for file in folder.data['files']:
if file.get('type') == 'file':
if Files.check_access_by_user_id(file.get('id'), user.id, 'read', db=db):
if await Files.check_access_by_user_id(file.get('id'), user.id, 'read', db=db):
valid_files.append(file)
elif file.get('type') == 'collection':
if Knowledges.check_access_by_user_id(file.get('id'), user.id, 'read', db=db):
if await Knowledges.check_access_by_user_id(file.get('id'), user.id, 'read', db=db):
valid_files.append(file)
else:
valid_files.append(file)
folder.data['files'] = valid_files
Folders.update_folder_by_id_and_user_id(folder.id, user.id, FolderUpdateForm(data=folder.data), db=db)
await Folders.update_folder_by_id_and_user_id(folder.id, user.id, FolderUpdateForm(data=folder.data), db=db)
folder_list.append(FolderNameIdResponse(**folder.model_dump()))
@ -102,12 +102,12 @@ async def get_folders(
@router.post('/')
def create_folder(
async def create_folder(
form_data: FolderForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
folder = Folders.get_folder_by_parent_id_and_user_id_and_name(form_data.parent_id, user.id, form_data.name, db=db)
folder = await Folders.get_folder_by_parent_id_and_user_id_and_name(form_data.parent_id, user.id, form_data.name, db=db)
if folder:
raise HTTPException(
@ -116,7 +116,7 @@ def create_folder(
)
try:
folder = Folders.insert_new_folder(user.id, form_data, form_data.parent_id, db=db)
folder = await Folders.insert_new_folder(user.id, form_data, form_data.parent_id, db=db)
return folder
except Exception as e:
log.exception(e)
@ -133,8 +133,8 @@ def create_folder(
@router.get('/{id}', response_model=Optional[FolderModel])
async def get_folder_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
folder = Folders.get_folder_by_id_and_user_id(id, user.id, db=db)
async def get_folder_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
folder = await Folders.get_folder_by_id_and_user_id(id, user.id, db=db)
if folder:
return folder
else:
@ -154,13 +154,13 @@ async def update_folder_name_by_id(
id: str,
form_data: FolderUpdateForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
folder = Folders.get_folder_by_id_and_user_id(id, user.id, db=db)
folder = await Folders.get_folder_by_id_and_user_id(id, user.id, db=db)
if folder:
if form_data.name is not None:
# Check if folder with same name exists
existing_folder = Folders.get_folder_by_parent_id_and_user_id_and_name(
existing_folder = await Folders.get_folder_by_parent_id_and_user_id_and_name(
folder.parent_id, user.id, form_data.name, db=db
)
if existing_folder and existing_folder.id != id:
@ -170,7 +170,7 @@ async def update_folder_name_by_id(
)
try:
folder = Folders.update_folder_by_id_and_user_id(id, user.id, form_data, db=db)
folder = await Folders.update_folder_by_id_and_user_id(id, user.id, form_data, db=db)
return folder
except Exception as e:
log.exception(e)
@ -200,11 +200,11 @@ async def update_folder_parent_id_by_id(
id: str,
form_data: FolderParentIdForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
folder = Folders.get_folder_by_id_and_user_id(id, user.id, db=db)
folder = await Folders.get_folder_by_id_and_user_id(id, user.id, db=db)
if folder:
existing_folder = Folders.get_folder_by_parent_id_and_user_id_and_name(
existing_folder = await Folders.get_folder_by_parent_id_and_user_id_and_name(
form_data.parent_id, user.id, folder.name, db=db
)
@ -215,7 +215,7 @@ async def update_folder_parent_id_by_id(
)
try:
folder = Folders.update_folder_parent_id_by_id_and_user_id(id, user.id, form_data.parent_id, db=db)
folder = await Folders.update_folder_parent_id_by_id_and_user_id(id, user.id, form_data.parent_id, db=db)
return folder
except Exception as e:
log.exception(e)
@ -245,12 +245,12 @@ async def update_folder_is_expanded_by_id(
id: str,
form_data: FolderIsExpandedForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
folder = Folders.get_folder_by_id_and_user_id(id, user.id, db=db)
folder = await Folders.get_folder_by_id_and_user_id(id, user.id, db=db)
if folder:
try:
folder = Folders.update_folder_is_expanded_by_id_and_user_id(id, user.id, form_data.is_expanded, db=db)
folder = await Folders.update_folder_is_expanded_by_id_and_user_id(id, user.id, form_data.is_expanded, db=db)
return folder
except Exception as e:
log.exception(e)
@ -277,10 +277,10 @@ async def delete_folder_by_id(
id: str,
delete_contents: Optional[bool] = True,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if Chats.count_chats_by_folder_id_and_user_id(id, user.id, db=db):
chat_delete_permission = has_permission(
if await Chats.count_chats_by_folder_id_and_user_id(id, user.id, db=db):
chat_delete_permission = await has_permission(
user.id, 'chat.delete', request.app.state.config.USER_PERMISSIONS, db=db
)
if user.role != 'admin' and not chat_delete_permission:
@ -290,18 +290,18 @@ async def delete_folder_by_id(
)
folders = []
folders.append(Folders.get_folder_by_id_and_user_id(id, user.id, db=db))
folders.append(await Folders.get_folder_by_id_and_user_id(id, user.id, db=db))
while folders:
folder = folders.pop()
if folder:
try:
folder_ids = Folders.delete_folder_by_id_and_user_id(folder.id, user.id, db=db)
folder_ids = await Folders.delete_folder_by_id_and_user_id(folder.id, user.id, db=db)
for folder_id in folder_ids:
if delete_contents:
Chats.delete_chats_by_user_id_and_folder_id(user.id, folder_id, db=db)
await Chats.delete_chats_by_user_id_and_folder_id(user.id, folder_id, db=db)
else:
Chats.move_chats_by_user_id_and_folder_id(user.id, folder_id, None, db=db)
await Chats.move_chats_by_user_id_and_folder_id(user.id, folder_id, None, db=db)
return True
except Exception as e:
@ -313,7 +313,7 @@ async def delete_folder_by_id(
)
finally:
# Get all subfolders
subfolders = Folders.get_folders_by_parent_id_and_user_id(folder.id, user.id, db=db)
subfolders = await Folders.get_folders_by_parent_id_and_user_id(folder.id, user.id, db=db)
folders.extend(subfolders)
else:

View file

@ -26,8 +26,8 @@ from open_webui.constants import ERROR_MESSAGES
from fastapi import APIRouter, Depends, HTTPException, Request, status
from open_webui.utils.auth import get_admin_user, get_verified_user
from pydantic import BaseModel, HttpUrl
from open_webui.internal.db import get_session
from sqlalchemy.orm import Session
from open_webui.internal.db import get_async_session
from sqlalchemy.ext.asyncio import AsyncSession
log = logging.getLogger(__name__)
@ -42,13 +42,13 @@ router = APIRouter()
@router.get('/', response_model=list[FunctionResponse])
async def get_functions(user=Depends(get_verified_user), db: Session = Depends(get_session)):
return Functions.get_functions(db=db)
async def get_functions(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
return await Functions.get_functions(db=db)
@router.get('/list', response_model=list[FunctionUserResponse])
async def get_function_list(user=Depends(get_admin_user), db: Session = Depends(get_session)):
return Functions.get_function_list(db=db)
async def get_function_list(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
return await Functions.get_function_list(db=db)
############################
@ -60,9 +60,9 @@ async def get_function_list(user=Depends(get_admin_user), db: Session = Depends(
async def get_functions(
include_valves: bool = False,
user=Depends(get_admin_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
return Functions.get_functions(include_valves=include_valves, db=db)
return await Functions.get_functions(include_valves=include_valves, db=db)
############################
@ -145,12 +145,12 @@ async def sync_functions(
request: Request,
form_data: SyncFunctionsForm,
user=Depends(get_admin_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
try:
for function in form_data.functions:
function.content = replace_imports(function.content)
function_module, function_type, frontmatter = load_function_module_by_id(
function_module, function_type, frontmatter = await load_function_module_by_id(
function.id,
content=function.content,
)
@ -163,7 +163,7 @@ async def sync_functions(
log.exception(f'Error validating valves for function {function.id}: {e}')
raise e
return Functions.sync_functions(user.id, form_data.functions, db=db)
return await Functions.sync_functions(user.id, form_data.functions, db=db)
except Exception as e:
log.exception(f'Failed to load a function: {e}')
raise HTTPException(
@ -182,7 +182,7 @@ async def create_new_function(
request: Request,
form_data: FunctionForm,
user=Depends(get_admin_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if not form_data.id.isidentifier():
raise HTTPException(
@ -192,11 +192,11 @@ async def create_new_function(
form_data.id = form_data.id.lower()
function = Functions.get_function_by_id(form_data.id, db=db)
function = await Functions.get_function_by_id(form_data.id, db=db)
if function is None:
try:
form_data.content = replace_imports(form_data.content)
function_module, function_type, frontmatter = load_function_module_by_id(
function_module, function_type, frontmatter = await load_function_module_by_id(
form_data.id,
content=form_data.content,
)
@ -205,13 +205,13 @@ async def create_new_function(
FUNCTIONS = request.app.state.FUNCTIONS
FUNCTIONS[form_data.id] = function_module
function = Functions.insert_new_function(user.id, function_type, form_data, db=db)
function = await Functions.insert_new_function(user.id, function_type, form_data, db=db)
function_cache_dir = CACHE_DIR / 'functions' / form_data.id
function_cache_dir.mkdir(parents=True, exist_ok=True)
if function_type == 'filter' and getattr(function_module, 'toggle', None):
Functions.update_function_metadata_by_id(form_data.id, {'toggle': True}, db=db)
await Functions.update_function_metadata_by_id(form_data.id, {'toggle': True}, db=db)
if function:
return function
@ -239,8 +239,8 @@ async def create_new_function(
@router.get('/id/{id}', response_model=Optional[FunctionModel])
async def get_function_by_id(id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)):
function = Functions.get_function_by_id(id, db=db)
async def get_function_by_id(id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
function = await Functions.get_function_by_id(id, db=db)
if function:
return function
@ -257,10 +257,10 @@ async def get_function_by_id(id: str, user=Depends(get_admin_user), db: Session
@router.post('/id/{id}/toggle', response_model=Optional[FunctionModel])
async def toggle_function_by_id(id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)):
function = Functions.get_function_by_id(id, db=db)
async def toggle_function_by_id(id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
function = await Functions.get_function_by_id(id, db=db)
if function:
function = Functions.update_function_by_id(id, {'is_active': not function.is_active}, db=db)
function = await Functions.update_function_by_id(id, {'is_active': not function.is_active}, db=db)
if function:
return function
@ -282,10 +282,10 @@ async def toggle_function_by_id(id: str, user=Depends(get_admin_user), db: Sessi
@router.post('/id/{id}/toggle/global', response_model=Optional[FunctionModel])
async def toggle_global_by_id(id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)):
function = Functions.get_function_by_id(id, db=db)
async def toggle_global_by_id(id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
function = await Functions.get_function_by_id(id, db=db)
if function:
function = Functions.update_function_by_id(id, {'is_global': not function.is_global}, db=db)
function = await Functions.update_function_by_id(id, {'is_global': not function.is_global}, db=db)
if function:
return function
@ -312,11 +312,11 @@ async def update_function_by_id(
id: str,
form_data: FunctionForm,
user=Depends(get_admin_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
try:
form_data.content = replace_imports(form_data.content)
function_module, function_type, frontmatter = load_function_module_by_id(id, content=form_data.content)
function_module, function_type, frontmatter = await load_function_module_by_id(id, content=form_data.content)
form_data.meta.manifest = frontmatter
FUNCTIONS = request.app.state.FUNCTIONS
@ -325,10 +325,10 @@ async def update_function_by_id(
updated = {**form_data.model_dump(exclude={'id'}), 'type': function_type}
log.debug(updated)
function = Functions.update_function_by_id(id, updated, db=db)
function = await Functions.update_function_by_id(id, updated, db=db)
if function_type == 'filter' and getattr(function_module, 'toggle', None):
Functions.update_function_metadata_by_id(id, {'toggle': True}, db=db)
await Functions.update_function_metadata_by_id(id, {'toggle': True}, db=db)
if function:
return function
@ -355,9 +355,9 @@ async def delete_function_by_id(
request: Request,
id: str,
user=Depends(get_admin_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
result = Functions.delete_function_by_id(id, db=db)
result = await Functions.delete_function_by_id(id, db=db)
if result:
FUNCTIONS = request.app.state.FUNCTIONS
@ -373,11 +373,11 @@ async def delete_function_by_id(
@router.get('/id/{id}/valves', response_model=Optional[dict])
async def get_function_valves_by_id(id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)):
function = Functions.get_function_by_id(id, db=db)
async def get_function_valves_by_id(id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
function = await Functions.get_function_by_id(id, db=db)
if function:
try:
valves = Functions.get_function_valves_by_id(id, db=db)
valves = await Functions.get_function_valves_by_id(id, db=db)
return valves
except Exception as e:
raise HTTPException(
@ -401,11 +401,11 @@ async def get_function_valves_spec_by_id(
request: Request,
id: str,
user=Depends(get_admin_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
function = Functions.get_function_by_id(id, db=db)
function = await Functions.get_function_by_id(id, db=db)
if function:
function_module, function_type, frontmatter = get_function_module_from_cache(request, id)
function_module, function_type, frontmatter = await get_function_module_from_cache(request, id)
if hasattr(function_module, 'Valves'):
Valves = function_module.Valves
@ -432,11 +432,11 @@ async def update_function_valves_by_id(
id: str,
form_data: dict,
user=Depends(get_admin_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
function = Functions.get_function_by_id(id, db=db)
function = await Functions.get_function_by_id(id, db=db)
if function:
function_module, function_type, frontmatter = get_function_module_from_cache(request, id)
function_module, function_type, frontmatter = await get_function_module_from_cache(request, id)
if hasattr(function_module, 'Valves'):
Valves = function_module.Valves
@ -446,7 +446,7 @@ async def update_function_valves_by_id(
valves = Valves(**form_data)
valves_dict = valves.model_dump(exclude_unset=True)
Functions.update_function_valves_by_id(id, valves_dict, db=db)
await Functions.update_function_valves_by_id(id, valves_dict, db=db)
return valves_dict
except Exception as e:
log.exception(f'Error updating function values by id {id}: {e}')
@ -473,11 +473,11 @@ async def update_function_valves_by_id(
@router.get('/id/{id}/valves/user', response_model=Optional[dict])
async def get_function_user_valves_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
function = Functions.get_function_by_id(id, db=db)
async def get_function_user_valves_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
function = await Functions.get_function_by_id(id, db=db)
if function:
try:
user_valves = Functions.get_user_valves_by_id_and_user_id(id, user.id, db=db)
user_valves = await Functions.get_user_valves_by_id_and_user_id(id, user.id, db=db)
return user_valves
except Exception as e:
raise HTTPException(
@ -496,11 +496,11 @@ async def get_function_user_valves_spec_by_id(
request: Request,
id: str,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
function = Functions.get_function_by_id(id, db=db)
function = await Functions.get_function_by_id(id, db=db)
if function:
function_module, function_type, frontmatter = get_function_module_from_cache(request, id)
function_module, function_type, frontmatter = await get_function_module_from_cache(request, id)
if hasattr(function_module, 'UserValves'):
UserValves = function_module.UserValves
@ -522,12 +522,12 @@ async def update_function_user_valves_by_id(
id: str,
form_data: dict,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
function = Functions.get_function_by_id(id, db=db)
function = await Functions.get_function_by_id(id, db=db)
if function:
function_module, function_type, frontmatter = get_function_module_from_cache(request, id)
function_module, function_type, frontmatter = await get_function_module_from_cache(request, id)
if hasattr(function_module, 'UserValves'):
UserValves = function_module.UserValves
@ -536,7 +536,7 @@ async def update_function_user_valves_by_id(
form_data = {k: v for k, v in form_data.items() if v is not None}
user_valves = UserValves(**form_data)
user_valves_dict = user_valves.model_dump(exclude_unset=True)
Functions.update_user_valves_by_id_and_user_id(id, user.id, user_valves_dict, db=db)
await Functions.update_user_valves_by_id_and_user_id(id, user.id, user_valves_dict, db=db)
return user_valves_dict
except Exception as e:
log.exception(f'Error updating function user valves by id {id}: {e}')

View file

@ -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:

View file

@ -28,8 +28,8 @@ from open_webui.routers.files import upload_file_handler, get_file_content_by_id
from open_webui.utils.auth import get_admin_user, get_verified_user
from open_webui.utils.access_control import has_permission
from open_webui.utils.headers import include_user_info_headers
from open_webui.internal.db import get_session
from sqlalchemy.orm import Session
from open_webui.internal.db import get_async_session
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.utils.images.comfyui import (
ComfyUICreateImageForm,
ComfyUIEditImageForm,
@ -341,7 +341,7 @@ async def verify_url(request: Request, user=Depends(get_admin_user)):
@router.get('/models')
def get_models(request: Request, user=Depends(get_verified_user)):
async def get_models(request: Request, user=Depends(get_verified_user)):
try:
if request.app.state.config.IMAGE_GENERATION_ENGINE == 'openai':
return [
@ -456,7 +456,7 @@ def get_image_data(data: str, headers=None):
return None, None
def upload_image(request, image_data, content_type, metadata, user, db=None):
async def upload_image(request, image_data, content_type, metadata, user, db=None):
image_format = mimetypes.guess_extension(content_type)
file = UploadFile(
file=io.BytesIO(image_data),
@ -465,7 +465,7 @@ def upload_image(request, image_data, content_type, metadata, user, db=None):
'content-type': content_type,
},
)
file_item = upload_file_handler(
file_item = await upload_file_handler(
request,
file=file,
metadata=metadata,
@ -479,7 +479,7 @@ def upload_image(request, image_data, content_type, metadata, user, db=None):
message_id = metadata.get('message_id')
if chat_id and message_id:
Chats.insert_chat_files(
await Chats.insert_chat_files(
chat_id=chat_id,
message_id=message_id,
file_ids=[file_item.id],
@ -499,7 +499,7 @@ async def generate_images(request: Request, form_data: CreateImageForm, user=Dep
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
if user.role != 'admin' and not has_permission(
if user.role != 'admin' and not await has_permission(
user.id, 'features.image_generation', request.app.state.config.USER_PERMISSIONS
):
raise HTTPException(
@ -590,7 +590,7 @@ async def image_generations(
else:
image_data, content_type = get_image_data(image['b64_json'])
_, url = upload_image(request, image_data, content_type, {**data, **metadata}, user)
_, url = await upload_image(request, image_data, content_type, {**data, **metadata}, user)
images.append({'url': url})
return images
@ -635,14 +635,14 @@ async def image_generations(
if model.endswith(':predict'):
for image in res['predictions']:
image_data, content_type = get_image_data(image['bytesBase64Encoded'])
_, url = upload_image(request, image_data, content_type, {**data, **metadata}, user)
_, url = await upload_image(request, image_data, content_type, {**data, **metadata}, user)
images.append({'url': url})
elif model.endswith(':generateContent'):
for image in res['candidates']:
for part in image['content']['parts']:
if part.get('inlineData', {}).get('data'):
image_data, content_type = get_image_data(part['inlineData']['data'])
_, url = upload_image(
_, url = await upload_image(
request,
image_data,
content_type,
@ -695,7 +695,7 @@ async def image_generations(
headers = {'Authorization': f'Bearer {request.app.state.config.COMFYUI_API_KEY}'}
image_data, content_type = get_image_data(image['url'], headers)
_, url = upload_image(
_, url = await upload_image(
request,
image_data,
content_type,
@ -742,7 +742,7 @@ async def image_generations(
for image in res['images']:
image_data, content_type = get_image_data(image)
_, url = upload_image(
_, url = await upload_image(
request,
image_data,
content_type,
@ -832,7 +832,7 @@ async def image_edits(
except Exception as e:
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.DEFAULT(e))
def get_image_file_item(base64_string, param_name='image'):
async def get_image_file_item(base64_string, param_name='image'):
data = base64_string
header, encoded = data.split(',', 1)
mime_type = header.split(';')[0].lstrip('data:')
@ -905,7 +905,7 @@ async def image_edits(
else:
image_data, content_type = get_image_data(image['b64_json'])
_, url = upload_image(request, image_data, content_type, {**data, **metadata}, user)
_, url = await upload_image(request, image_data, content_type, {**data, **metadata}, user)
images.append({'url': url})
return images
@ -956,7 +956,7 @@ async def image_edits(
for part in image['content']['parts']:
if part.get('inlineData', {}).get('data'):
image_data, content_type = get_image_data(part['inlineData']['data'])
_, url = upload_image(
_, url = await upload_image(
request,
image_data,
content_type,
@ -1036,7 +1036,7 @@ async def image_edits(
headers = {'Authorization': f'Bearer {request.app.state.config.IMAGES_EDIT_COMFYUI_API_KEY}'}
image_data, content_type = get_image_data(image_url, headers)
_, url = upload_image(
_, url = await upload_image(
request,
image_data,
content_type,

View file

@ -8,8 +8,8 @@ import io
import zipfile
from urllib.parse import quote
from sqlalchemy.orm import Session
from open_webui.internal.db import get_session
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import get_async_session
from open_webui.models.groups import Groups
from open_webui.models.knowledge import (
KnowledgeFileListResponse,
@ -111,14 +111,14 @@ class KnowledgeAccessListResponse(BaseModel):
async def get_knowledge_bases(
page: Optional[int] = 1,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
page = max(page, 1)
limit = PAGE_ITEM_COUNT
skip = (page - 1) * limit
filter = {}
groups = Groups.get_groups_by_member_id(user.id, db=db)
groups = await Groups.get_groups_by_member_id(user.id, db=db)
user_group_ids = {group.id for group in groups}
if not user.role == 'admin' or not BYPASS_ADMIN_ACCESS_CONTROL:
@ -127,11 +127,11 @@ async def get_knowledge_bases(
filter['user_id'] = user.id
result = Knowledges.search_knowledge_bases(user.id, filter=filter, skip=skip, limit=limit, db=db)
result = await Knowledges.search_knowledge_bases(user.id, filter=filter, skip=skip, limit=limit, db=db)
# Batch-fetch writable knowledge IDs in a single query instead of N has_access calls
knowledge_base_ids = [knowledge_base.id for knowledge_base in result.items]
writable_knowledge_base_ids = AccessGrants.get_accessible_resource_ids(
writable_knowledge_base_ids = await AccessGrants.get_accessible_resource_ids(
user_id=user.id,
resource_type='knowledge',
resource_ids=knowledge_base_ids,
@ -162,7 +162,7 @@ async def search_knowledge_bases(
view_option: Optional[str] = None,
page: Optional[int] = 1,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
page = max(page, 1)
limit = PAGE_ITEM_COUNT
@ -174,7 +174,7 @@ async def search_knowledge_bases(
if view_option:
filter['view_option'] = view_option
groups = Groups.get_groups_by_member_id(user.id, db=db)
groups = await Groups.get_groups_by_member_id(user.id, db=db)
user_group_ids = {group.id for group in groups}
if not user.role == 'admin' or not BYPASS_ADMIN_ACCESS_CONTROL:
@ -183,11 +183,11 @@ async def search_knowledge_bases(
filter['user_id'] = user.id
result = Knowledges.search_knowledge_bases(user.id, filter=filter, skip=skip, limit=limit, db=db)
result = await Knowledges.search_knowledge_bases(user.id, filter=filter, skip=skip, limit=limit, db=db)
# Batch-fetch writable knowledge IDs in a single query instead of N has_access calls
knowledge_base_ids = [knowledge_base.id for knowledge_base in result.items]
writable_knowledge_base_ids = AccessGrants.get_accessible_resource_ids(
writable_knowledge_base_ids = await AccessGrants.get_accessible_resource_ids(
user_id=user.id,
resource_type='knowledge',
resource_ids=knowledge_base_ids,
@ -217,7 +217,7 @@ async def search_knowledge_files(
query: Optional[str] = None,
page: Optional[int] = 1,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
page = max(page, 1)
limit = PAGE_ITEM_COUNT
@ -227,13 +227,13 @@ async def search_knowledge_files(
if query:
filter['query'] = query
groups = Groups.get_groups_by_member_id(user.id, db=db)
groups = await Groups.get_groups_by_member_id(user.id, db=db)
if groups:
filter['group_ids'] = [group.id for group in groups]
filter['user_id'] = user.id
return Knowledges.search_knowledge_files(filter=filter, skip=skip, limit=limit, db=db)
return await Knowledges.search_knowledge_files(filter=filter, skip=skip, limit=limit, db=db)
############################
@ -247,11 +247,11 @@ async def create_new_knowledge(
form_data: KnowledgeForm,
user=Depends(get_verified_user),
):
# NOTE: We intentionally do NOT use Depends(get_session) here.
# NOTE: We intentionally do NOT use Depends(get_async_session) here.
# Database operations (has_permission, filter_allowed_access_grants, insert_new_knowledge) manage their own sessions.
# This prevents holding a connection during embed_knowledge_base_metadata()
# which makes external embedding API calls (1-5+ seconds).
if user.role != 'admin' and not has_permission(
if user.role != 'admin' and not await has_permission(
user.id, 'workspace.knowledge', request.app.state.config.USER_PERMISSIONS
):
raise HTTPException(
@ -259,7 +259,7 @@ async def create_new_knowledge(
detail=ERROR_MESSAGES.UNAUTHORIZED,
)
form_data.access_grants = filter_allowed_access_grants(
form_data.access_grants = await filter_allowed_access_grants(
request.app.state.config.USER_PERMISSIONS,
user.id,
user.role,
@ -267,7 +267,7 @@ async def create_new_knowledge(
'sharing.public_knowledge',
)
knowledge = Knowledges.insert_new_knowledge(user.id, form_data)
knowledge = await Knowledges.insert_new_knowledge(user.id, form_data)
if knowledge:
# Embed knowledge base for semantic search
@ -294,7 +294,7 @@ async def create_new_knowledge(
async def reindex_knowledge_files(
request: Request,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if user.role != 'admin':
raise HTTPException(
@ -302,13 +302,13 @@ async def reindex_knowledge_files(
detail=ERROR_MESSAGES.UNAUTHORIZED,
)
knowledge_bases = Knowledges.get_knowledge_bases(db=db)
knowledge_bases = await Knowledges.get_knowledge_bases(db=db)
log.info(f'Starting reindexing for {len(knowledge_bases)} knowledge bases')
for knowledge_base in knowledge_bases:
try:
files = Knowledges.get_files_by_id(knowledge_base.id, db=db)
files = await Knowledges.get_files_by_id(knowledge_base.id, db=db)
try:
if VECTOR_DB_CLIENT.has_collection(collection_name=knowledge_base.id):
VECTOR_DB_CLIENT.delete_collection(collection_name=knowledge_base.id)
@ -357,12 +357,12 @@ async def reindex_knowledge_base_metadata_embeddings(
):
"""Batch embed all existing knowledge bases. Admin only.
NOTE: We intentionally do NOT use Depends(get_session) here.
NOTE: We intentionally do NOT use Depends(get_async_session) here.
This endpoint loops through ALL knowledge bases and calls embed_knowledge_base_metadata()
for each one, making N external embedding API calls. Holding a session during
this entire operation would exhaust the connection pool.
"""
knowledge_bases = Knowledges.get_knowledge_bases()
knowledge_bases = await Knowledges.get_knowledge_bases()
log.info(f'Reindexing embeddings for {len(knowledge_bases)} knowledge bases')
success_count = 0
@ -385,14 +385,14 @@ class KnowledgeFilesResponse(KnowledgeResponse):
@router.get('/{id}', response_model=Optional[KnowledgeFilesResponse])
async def get_knowledge_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
knowledge = Knowledges.get_knowledge_by_id(id=id, db=db)
async def get_knowledge_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db)
if knowledge:
if (
user.role == 'admin'
or knowledge.user_id == user.id
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user.id,
resource_type='knowledge',
resource_id=knowledge.id,
@ -405,7 +405,7 @@ async def get_knowledge_by_id(id: str, user=Depends(get_verified_user), db: Sess
write_access=(
user.id == knowledge.user_id
or (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL)
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user.id,
resource_type='knowledge',
resource_id=knowledge.id,
@ -438,11 +438,11 @@ async def update_knowledge_by_id(
form_data: KnowledgeForm,
user=Depends(get_verified_user),
):
# NOTE: We intentionally do NOT use Depends(get_session) here.
# NOTE: We intentionally do NOT use Depends(get_async_session) here.
# Database operations manage their own short-lived sessions internally.
# This prevents holding a connection during embed_knowledge_base_metadata()
# which makes external embedding API calls (1-5+ seconds).
knowledge = Knowledges.get_knowledge_by_id(id=id)
knowledge = await Knowledges.get_knowledge_by_id(id=id)
if not knowledge:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
@ -451,7 +451,7 @@ async def update_knowledge_by_id(
# Is the user the original creator, in a group with write access, or an admin
if (
knowledge.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='knowledge',
resource_id=knowledge.id,
@ -464,7 +464,7 @@ async def update_knowledge_by_id(
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
form_data.access_grants = filter_allowed_access_grants(
form_data.access_grants = await filter_allowed_access_grants(
request.app.state.config.USER_PERMISSIONS,
user.id,
user.role,
@ -472,7 +472,7 @@ async def update_knowledge_by_id(
'sharing.public_knowledge',
)
knowledge = Knowledges.update_knowledge_by_id(id=id, form_data=form_data)
knowledge = await Knowledges.update_knowledge_by_id(id=id, form_data=form_data)
if knowledge:
# Re-embed knowledge base for semantic search
await embed_knowledge_base_metadata(
@ -483,7 +483,7 @@ async def update_knowledge_by_id(
)
return KnowledgeFilesResponse(
**knowledge.model_dump(),
files=Knowledges.get_file_metadatas_by_id(knowledge.id),
files=await Knowledges.get_file_metadatas_by_id(knowledge.id),
)
else:
raise HTTPException(
@ -507,9 +507,9 @@ async def update_knowledge_access_by_id(
id: str,
form_data: KnowledgeAccessGrantsForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
knowledge = Knowledges.get_knowledge_by_id(id=id, db=db)
knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db)
if not knowledge:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
@ -518,7 +518,7 @@ async def update_knowledge_access_by_id(
if (
knowledge.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='knowledge',
resource_id=knowledge.id,
@ -532,7 +532,7 @@ async def update_knowledge_access_by_id(
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
form_data.access_grants = filter_allowed_access_grants(
form_data.access_grants = await filter_allowed_access_grants(
request.app.state.config.USER_PERMISSIONS,
user.id,
user.role,
@ -540,11 +540,11 @@ async def update_knowledge_access_by_id(
'sharing.public_knowledge',
)
AccessGrants.set_access_grants('knowledge', id, form_data.access_grants, db=db)
await AccessGrants.set_access_grants('knowledge', id, form_data.access_grants, db=db)
return KnowledgeFilesResponse(
**Knowledges.get_knowledge_by_id(id=id, db=db).model_dump(),
files=Knowledges.get_file_metadatas_by_id(id, db=db),
**await Knowledges.get_knowledge_by_id(id=id, db=db).model_dump(),
files=await Knowledges.get_file_metadatas_by_id(id, db=db),
)
@ -562,9 +562,9 @@ async def get_knowledge_files_by_id(
direction: Optional[str] = None,
page: Optional[int] = 1,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
knowledge = Knowledges.get_knowledge_by_id(id=id, db=db)
knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db)
if not knowledge:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
@ -574,7 +574,7 @@ async def get_knowledge_files_by_id(
if not (
user.role == 'admin'
or knowledge.user_id == user.id
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user.id,
resource_type='knowledge',
resource_id=knowledge.id,
@ -602,7 +602,7 @@ async def get_knowledge_files_by_id(
if direction:
filter['direction'] = direction
return Knowledges.search_files_by_id(id, user.id, filter=filter, skip=skip, limit=limit, db=db)
return await Knowledges.search_files_by_id(id, user.id, filter=filter, skip=skip, limit=limit, db=db)
############################
@ -615,14 +615,14 @@ class KnowledgeFileIdForm(BaseModel):
@router.post('/{id}/file/add', response_model=Optional[KnowledgeFilesResponse])
def add_file_to_knowledge_by_id(
async def add_file_to_knowledge_by_id(
request: Request,
id: str,
form_data: KnowledgeFileIdForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
knowledge = Knowledges.get_knowledge_by_id(id=id, db=db)
knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db)
if not knowledge:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
@ -631,7 +631,7 @@ def add_file_to_knowledge_by_id(
if (
knowledge.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='knowledge',
resource_id=knowledge.id,
@ -645,7 +645,7 @@ def add_file_to_knowledge_by_id(
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
file = Files.get_file_by_id(form_data.file_id, db=db)
file = await Files.get_file_by_id(form_data.file_id, db=db)
if not file:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
@ -667,7 +667,7 @@ def add_file_to_knowledge_by_id(
)
# Add file to knowledge base
Knowledges.add_file_to_knowledge_by_id(knowledge_id=id, file_id=form_data.file_id, user_id=user.id, db=db)
await Knowledges.add_file_to_knowledge_by_id(knowledge_id=id, file_id=form_data.file_id, user_id=user.id, db=db)
except Exception as e:
log.debug(e)
raise HTTPException(
@ -678,7 +678,7 @@ def add_file_to_knowledge_by_id(
if knowledge:
return KnowledgeFilesResponse(
**knowledge.model_dump(),
files=Knowledges.get_file_metadatas_by_id(knowledge.id, db=db),
files=await Knowledges.get_file_metadatas_by_id(knowledge.id, db=db),
)
else:
raise HTTPException(
@ -688,14 +688,14 @@ def add_file_to_knowledge_by_id(
@router.post('/{id}/file/update', response_model=Optional[KnowledgeFilesResponse])
def update_file_from_knowledge_by_id(
async def update_file_from_knowledge_by_id(
request: Request,
id: str,
form_data: KnowledgeFileIdForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
knowledge = Knowledges.get_knowledge_by_id(id=id, db=db)
knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db)
if not knowledge:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
@ -704,7 +704,7 @@ def update_file_from_knowledge_by_id(
if (
knowledge.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='knowledge',
resource_id=knowledge.id,
@ -718,7 +718,7 @@ def update_file_from_knowledge_by_id(
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
file = Files.get_file_by_id(form_data.file_id, db=db)
file = await Files.get_file_by_id(form_data.file_id, db=db)
if not file:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
@ -726,7 +726,7 @@ def update_file_from_knowledge_by_id(
)
# Validate the file actually belongs to this knowledge base
if not Knowledges.has_file(knowledge_id=id, file_id=form_data.file_id, db=db):
if not await Knowledges.has_file(knowledge_id=id, file_id=form_data.file_id, db=db):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=ERROR_MESSAGES.NOT_FOUND,
@ -752,7 +752,7 @@ def update_file_from_knowledge_by_id(
if knowledge:
return KnowledgeFilesResponse(
**knowledge.model_dump(),
files=Knowledges.get_file_metadatas_by_id(knowledge.id, db=db),
files=await Knowledges.get_file_metadatas_by_id(knowledge.id, db=db),
)
else:
raise HTTPException(
@ -767,14 +767,14 @@ def update_file_from_knowledge_by_id(
@router.post('/{id}/file/remove', response_model=Optional[KnowledgeFilesResponse])
def remove_file_from_knowledge_by_id(
async def remove_file_from_knowledge_by_id(
id: str,
form_data: KnowledgeFileIdForm,
delete_file: bool = Query(True),
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
knowledge = Knowledges.get_knowledge_by_id(id=id, db=db)
knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db)
if not knowledge:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
@ -783,7 +783,7 @@ def remove_file_from_knowledge_by_id(
if (
knowledge.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='knowledge',
resource_id=knowledge.id,
@ -797,7 +797,7 @@ def remove_file_from_knowledge_by_id(
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
file = Files.get_file_by_id(form_data.file_id, db=db)
file = await Files.get_file_by_id(form_data.file_id, db=db)
if not file:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
@ -805,13 +805,13 @@ def remove_file_from_knowledge_by_id(
)
# Validate the file actually belongs to this knowledge base
if not Knowledges.has_file(knowledge_id=id, file_id=form_data.file_id, db=db):
if not await Knowledges.has_file(knowledge_id=id, file_id=form_data.file_id, db=db):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=ERROR_MESSAGES.NOT_FOUND,
)
Knowledges.remove_file_from_knowledge_by_id(knowledge_id=id, file_id=form_data.file_id, db=db)
await Knowledges.remove_file_from_knowledge_by_id(knowledge_id=id, file_id=form_data.file_id, db=db)
# Remove content from the vector database
try:
@ -839,12 +839,12 @@ def remove_file_from_knowledge_by_id(
pass
# Delete file from database
Files.delete_file_by_id(form_data.file_id, db=db)
await Files.delete_file_by_id(form_data.file_id, db=db)
if knowledge:
return KnowledgeFilesResponse(
**knowledge.model_dump(),
files=Knowledges.get_file_metadatas_by_id(knowledge.id, db=db),
files=await Knowledges.get_file_metadatas_by_id(knowledge.id, db=db),
)
else:
raise HTTPException(
@ -859,8 +859,8 @@ def remove_file_from_knowledge_by_id(
@router.delete('/{id}/delete', response_model=bool)
async def delete_knowledge_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
knowledge = Knowledges.get_knowledge_by_id(id=id, db=db)
async def delete_knowledge_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db)
if not knowledge:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
@ -869,7 +869,7 @@ async def delete_knowledge_by_id(id: str, user=Depends(get_verified_user), db: S
if (
knowledge.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='knowledge',
resource_id=knowledge.id,
@ -886,7 +886,7 @@ async def delete_knowledge_by_id(id: str, user=Depends(get_verified_user), db: S
log.info(f'Deleting knowledge base: {id} (name: {knowledge.name})')
# Get all models
models = Models.get_all_models(db=db)
models = await Models.get_all_models(db=db)
log.info(f'Found {len(models)} models to check for knowledge base {id}')
# Update models that reference this knowledge base
@ -910,7 +910,7 @@ async def delete_knowledge_by_id(id: str, user=Depends(get_verified_user), db: S
access_grants=model.access_grants,
is_active=model.is_active,
)
Models.update_model_by_id(model.id, model_form, db=db)
await Models.update_model_by_id(model.id, model_form, db=db)
# Clean up vector DB
try:
@ -922,7 +922,7 @@ async def delete_knowledge_by_id(id: str, user=Depends(get_verified_user), db: S
# Remove knowledge base embedding
remove_knowledge_base_metadata_embedding(id)
result = Knowledges.delete_knowledge_by_id(id=id, db=db)
result = await Knowledges.delete_knowledge_by_id(id=id, db=db)
return result
@ -932,8 +932,8 @@ async def delete_knowledge_by_id(id: str, user=Depends(get_verified_user), db: S
@router.post('/{id}/reset', response_model=Optional[KnowledgeResponse])
async def reset_knowledge_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
knowledge = Knowledges.get_knowledge_by_id(id=id, db=db)
async def reset_knowledge_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db)
if not knowledge:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
@ -942,7 +942,7 @@ async def reset_knowledge_by_id(id: str, user=Depends(get_verified_user), db: Se
if (
knowledge.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='knowledge',
resource_id=knowledge.id,
@ -977,12 +977,12 @@ async def add_files_to_knowledge_batch(
id: str,
form_data: list[KnowledgeFileIdForm],
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""
Add multiple files to a knowledge base
"""
knowledge = Knowledges.get_knowledge_by_id(id=id, db=db)
knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db)
if not knowledge:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
@ -991,7 +991,7 @@ async def add_files_to_knowledge_batch(
if (
knowledge.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='knowledge',
resource_id=knowledge.id,
@ -1008,7 +1008,7 @@ async def add_files_to_knowledge_batch(
# Batch-fetch all files to avoid N+1 queries
log.info(f'files/batch/add - {len(form_data)} files')
file_ids = [form.file_id for form in form_data]
files = Files.get_files_by_ids(file_ids, db=db)
files = await Files.get_files_by_ids(file_ids, db=db)
# Verify all requested files were found
found_ids = {file.id for file in files}
@ -1034,14 +1034,14 @@ async def add_files_to_knowledge_batch(
# Only add files that were successfully processed
successful_file_ids = [r.file_id for r in result.results if r.status == 'completed']
for file_id in successful_file_ids:
Knowledges.add_file_to_knowledge_by_id(knowledge_id=id, file_id=file_id, user_id=user.id, db=db)
await Knowledges.add_file_to_knowledge_by_id(knowledge_id=id, file_id=file_id, user_id=user.id, db=db)
# If there were any errors, include them in the response
if result.errors:
error_details = [f'{err.file_id}: {err.error}' for err in result.errors]
return KnowledgeFilesResponse(
**knowledge.model_dump(),
files=Knowledges.get_file_metadatas_by_id(knowledge.id, db=db),
files=await Knowledges.get_file_metadatas_by_id(knowledge.id, db=db),
warnings={
'message': 'Some files failed to process',
'errors': error_details,
@ -1050,7 +1050,7 @@ async def add_files_to_knowledge_batch(
return KnowledgeFilesResponse(
**knowledge.model_dump(),
files=Knowledges.get_file_metadatas_by_id(knowledge.id, db=db),
files=await Knowledges.get_file_metadatas_by_id(knowledge.id, db=db),
)
@ -1060,20 +1060,20 @@ async def add_files_to_knowledge_batch(
@router.get('/{id}/export')
async def export_knowledge_by_id(id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)):
async def export_knowledge_by_id(id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
"""
Export a knowledge base as a zip file containing .txt files.
Admin only.
"""
knowledge = Knowledges.get_knowledge_by_id(id=id, db=db)
knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db)
if not knowledge:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.NOT_FOUND,
)
files = Knowledges.get_files_by_id(id, db=db)
files = await Knowledges.get_files_by_id(id, db=db)
# Create zip file in memory
zip_buffer = io.BytesIO()

View file

@ -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])

View file

@ -35,8 +35,8 @@ from fastapi.responses import FileResponse, StreamingResponse
from open_webui.utils.auth import get_admin_user, get_verified_user
from open_webui.utils.access_control import has_permission, filter_allowed_access_grants
from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL, STATIC_DIR
from open_webui.internal.db import get_session
from sqlalchemy.orm import Session
from open_webui.internal.db import get_async_session
from sqlalchemy.ext.asyncio import AsyncSession
log = logging.getLogger(__name__)
@ -66,7 +66,7 @@ async def get_models(
direction: Optional[str] = None,
page: Optional[int] = 1,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
limit = PAGE_ITEM_COUNT
@ -86,7 +86,7 @@ async def get_models(
filter['direction'] = direction
# Pre-fetch user group IDs once - used for both filter and write_access check
groups = Groups.get_groups_by_member_id(user.id, db=db)
groups = await Groups.get_groups_by_member_id(user.id, db=db)
user_group_ids = {group.id for group in groups}
if not user.role == 'admin' or not BYPASS_ADMIN_ACCESS_CONTROL:
@ -95,11 +95,11 @@ async def get_models(
filter['user_id'] = user.id
result = Models.search_models(user.id, filter=filter, skip=skip, limit=limit, db=db)
result = await Models.search_models(user.id, filter=filter, skip=skip, limit=limit, db=db)
# Batch-fetch writable model IDs in a single query instead of N has_access calls
model_ids = [model.id for model in result.items]
writable_model_ids = AccessGrants.get_accessible_resource_ids(
writable_model_ids = await AccessGrants.get_accessible_resource_ids(
user_id=user.id,
resource_type='model',
resource_ids=model_ids,
@ -130,8 +130,8 @@ async def get_models(
@router.get('/base', response_model=list[ModelResponse])
async def get_base_models(user=Depends(get_admin_user), db: Session = Depends(get_session)):
return Models.get_base_models(db=db)
async def get_base_models(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
return await Models.get_base_models(db=db)
###########################
@ -140,11 +140,11 @@ async def get_base_models(user=Depends(get_admin_user), db: Session = Depends(ge
@router.get('/tags', response_model=list[str])
async def get_model_tags(user=Depends(get_verified_user), db: Session = Depends(get_session)):
async def get_model_tags(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
if user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL:
models = Models.get_models(db=db)
models = await Models.get_models(db=db)
else:
models = Models.get_models_by_user_id(user.id, db=db)
models = await Models.get_models_by_user_id(user.id, db=db)
tags_set = set()
for model in models:
@ -172,9 +172,9 @@ async def create_new_model(
request: Request,
form_data: ModelForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if user.role != 'admin' and not has_permission(
if user.role != 'admin' and not await has_permission(
user.id, 'workspace.models', request.app.state.config.USER_PERMISSIONS, db=db
):
raise HTTPException(
@ -182,7 +182,7 @@ async def create_new_model(
detail=ERROR_MESSAGES.UNAUTHORIZED,
)
model = Models.get_model_by_id(form_data.id, db=db)
model = await Models.get_model_by_id(form_data.id, db=db)
if model:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
@ -196,7 +196,7 @@ async def create_new_model(
)
else:
form_data.access_grants = filter_allowed_access_grants(
form_data.access_grants = await filter_allowed_access_grants(
request.app.state.config.USER_PERMISSIONS,
user.id,
user.role,
@ -204,7 +204,7 @@ async def create_new_model(
'sharing.public_models',
)
model = Models.insert_new_model(form_data, user.id, db=db)
model = await Models.insert_new_model(form_data, user.id, db=db)
if model:
return model
else:
@ -223,9 +223,9 @@ async def create_new_model(
async def export_models(
request: Request,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if user.role != 'admin' and not has_permission(
if user.role != 'admin' and not await has_permission(
user.id,
'workspace.models_export',
request.app.state.config.USER_PERMISSIONS,
@ -237,9 +237,9 @@ async def export_models(
)
if user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL:
return Models.get_models(db=db)
return await Models.get_models(db=db)
else:
return Models.get_models_by_user_id(user.id, db=db)
return await Models.get_models_by_user_id(user.id, db=db)
############################
@ -256,9 +256,9 @@ async def import_models(
request: Request,
user=Depends(get_verified_user),
form_data: ModelsImportForm = (...),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if user.role != 'admin' and not has_permission(
if user.role != 'admin' and not await has_permission(
user.id,
'workspace.models_import',
request.app.state.config.USER_PERMISSIONS,
@ -278,7 +278,7 @@ async def import_models(
if model_data.get('id') and is_valid_model_id(model_data.get('id'))
]
existing_models = {
model.id: model for model in (Models.get_models_by_ids(model_ids, db=db) if model_ids else [])
model.id: model for model in (await Models.get_models_by_ids(model_ids, db=db) if model_ids else [])
}
for model_data in data:
@ -293,13 +293,13 @@ async def import_models(
model_data['params'] = model_data.get('params', {})
updated_model = ModelForm(**{**existing_model.model_dump(), **model_data})
Models.update_model_by_id(model_id, updated_model, db=db)
await Models.update_model_by_id(model_id, updated_model, db=db)
else:
# Insert new model
model_data['meta'] = model_data.get('meta', {})
model_data['params'] = model_data.get('params', {})
new_model = ModelForm(**model_data)
Models.insert_new_model(user_id=user.id, form_data=new_model, db=db)
await Models.insert_new_model(user_id=user.id, form_data=new_model, db=db)
return True
else:
raise HTTPException(status_code=400, detail='Invalid JSON format')
@ -322,9 +322,9 @@ async def sync_models(
request: Request,
form_data: SyncModelsForm,
user=Depends(get_admin_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
return Models.sync_models(user.id, form_data.models, db=db)
return await Models.sync_models(user.id, form_data.models, db=db)
###########################
@ -338,13 +338,13 @@ class ModelIdForm(BaseModel):
# Note: We're not using the typical url path param here, but instead using a query parameter to allow '/' in the id
@router.get('/model', response_model=Optional[ModelAccessResponse])
async def get_model_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
model = Models.get_model_by_id(id, db=db)
async def get_model_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
model = await Models.get_model_by_id(id, db=db)
if model:
if (
(user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL)
or model.user_id == user.id
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user.id,
resource_type='model',
resource_id=model.id,
@ -357,7 +357,7 @@ async def get_model_by_id(id: str, user=Depends(get_verified_user), db: Session
write_access=(
(user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL)
or user.id == model.user_id
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user.id,
resource_type='model',
resource_id=model.id,
@ -384,8 +384,8 @@ async def get_model_by_id(id: str, user=Depends(get_verified_user), db: Session
@router.get('/model/profile/image')
def get_model_profile_image(id: str, user=Depends(get_verified_user)):
model = Models.get_model_by_id(id)
async def get_model_profile_image(id: str, user=Depends(get_verified_user)):
model = await Models.get_model_by_id(id)
if model:
etag = f'"{model.updated_at}"' if model.updated_at else None
@ -426,13 +426,13 @@ def get_model_profile_image(id: str, user=Depends(get_verified_user)):
@router.post('/model/toggle', response_model=Optional[ModelResponse])
async def toggle_model_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
model = Models.get_model_by_id(id, db=db)
async def toggle_model_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
model = await Models.get_model_by_id(id, db=db)
if model:
if (
user.role == 'admin'
or model.user_id == user.id
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user.id,
resource_type='model',
resource_id=model.id,
@ -440,7 +440,7 @@ async def toggle_model_by_id(id: str, user=Depends(get_verified_user), db: Sessi
db=db,
)
):
model = Models.toggle_model_by_id(id, db=db)
model = await Models.toggle_model_by_id(id, db=db)
if model:
return model
@ -471,9 +471,9 @@ async def update_model_by_id(
request: Request,
form_data: ModelForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
model = Models.get_model_by_id(form_data.id, db=db)
model = await Models.get_model_by_id(form_data.id, db=db)
if not model:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
@ -482,7 +482,7 @@ async def update_model_by_id(
if (
model.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='model',
resource_id=model.id,
@ -496,7 +496,7 @@ async def update_model_by_id(
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
form_data.access_grants = filter_allowed_access_grants(
form_data.access_grants = await filter_allowed_access_grants(
request.app.state.config.USER_PERMISSIONS,
user.id,
user.role,
@ -504,7 +504,7 @@ async def update_model_by_id(
'sharing.public_models',
)
model = Models.update_model_by_id(form_data.id, ModelForm(**form_data.model_dump()), db=db)
model = await Models.update_model_by_id(form_data.id, ModelForm(**form_data.model_dump()), db=db)
return model
@ -524,9 +524,9 @@ async def update_model_access_by_id(
request: Request,
form_data: ModelAccessGrantsForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
model = Models.get_model_by_id(form_data.id, db=db)
model = await Models.get_model_by_id(form_data.id, db=db)
# Non-preset models (e.g. direct Ollama/OpenAI models) may not have a DB
# entry yet. Create a minimal one so access grants can be stored.
@ -536,7 +536,7 @@ async def update_model_access_by_id(
status_code=status.HTTP_403_FORBIDDEN,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
model = Models.insert_new_model(
model = await Models.insert_new_model(
ModelForm(
id=form_data.id,
name=form_data.name or form_data.id,
@ -554,7 +554,7 @@ async def update_model_access_by_id(
if (
model.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='model',
resource_id=model.id,
@ -568,7 +568,7 @@ async def update_model_access_by_id(
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
form_data.access_grants = filter_allowed_access_grants(
form_data.access_grants = await filter_allowed_access_grants(
request.app.state.config.USER_PERMISSIONS,
user.id,
user.role,
@ -576,11 +576,11 @@ async def update_model_access_by_id(
'sharing.public_models',
)
AccessGrants.set_access_grants('model', form_data.id, form_data.access_grants, db=db)
await AccessGrants.set_access_grants('model', form_data.id, form_data.access_grants, db=db)
Models.update_model_updated_at_by_id(form_data.id, db=db)
await Models.update_model_updated_at_by_id(form_data.id, db=db)
return Models.get_model_by_id(form_data.id, db=db)
return await Models.get_model_by_id(form_data.id, db=db)
############################
@ -592,9 +592,9 @@ async def update_model_access_by_id(
async def delete_model_by_id(
form_data: ModelIdForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
model = Models.get_model_by_id(form_data.id, db=db)
model = await Models.get_model_by_id(form_data.id, db=db)
if not model:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
@ -604,7 +604,7 @@ async def delete_model_by_id(
if (
user.role != 'admin'
and model.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='model',
resource_id=model.id,
@ -617,11 +617,11 @@ async def delete_model_by_id(
detail=ERROR_MESSAGES.UNAUTHORIZED,
)
result = Models.delete_model_by_id(form_data.id, db=db)
result = await Models.delete_model_by_id(form_data.id, db=db)
return result
@router.delete('/delete/all', response_model=bool)
async def delete_all_models(user=Depends(get_admin_user), db: Session = Depends(get_session)):
result = Models.delete_all_models(db=db)
async def delete_all_models(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
result = await Models.delete_all_models(db=db)
return result

View file

@ -34,8 +34,8 @@ from open_webui.utils.access_control import (
filter_allowed_access_grants,
)
from open_webui.models.access_grants import AccessGrants
from open_webui.internal.db import get_session
from sqlalchemy.orm import Session
from open_webui.internal.db import get_async_session
from sqlalchemy.ext.asyncio import AsyncSession
log = logging.getLogger(__name__)
@ -68,9 +68,9 @@ async def get_notes(
request: Request,
page: Optional[int] = None,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if user.role != 'admin' and not has_permission(
if user.role != 'admin' and not await has_permission(
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
):
raise HTTPException(
@ -84,12 +84,12 @@ async def get_notes(
limit = 60
skip = (page - 1) * limit
notes = Notes.get_notes_by_user_id(user.id, 'read', skip=skip, limit=limit, db=db)
notes = await Notes.get_notes_by_user_id(user.id, 'read', skip=skip, limit=limit, db=db)
if not notes:
return []
user_ids = list(set(note.user_id for note in notes))
users = {user.id: user for user in Users.get_users_by_user_ids(user_ids, db=db)}
users = {user.id: user for user in await Users.get_users_by_user_ids(user_ids, db=db)}
return [
NoteUserResponse(
@ -114,9 +114,9 @@ async def search_notes(
direction: Optional[str] = None,
page: Optional[int] = 1,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if user.role != 'admin' and not has_permission(
if user.role != 'admin' and not await has_permission(
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
):
raise HTTPException(
@ -143,13 +143,13 @@ async def search_notes(
filter['direction'] = direction
if not user.role == 'admin' or not BYPASS_ADMIN_ACCESS_CONTROL:
groups = Groups.get_groups_by_member_id(user.id, db=db)
groups = await Groups.get_groups_by_member_id(user.id, db=db)
if groups:
filter['group_ids'] = [group.id for group in groups]
filter['user_id'] = user.id
result = Notes.search_notes(user.id, filter, skip=skip, limit=limit, db=db)
result = await Notes.search_notes(user.id, filter, skip=skip, limit=limit, db=db)
for note in result.items:
note.data = _truncate_note_data(note.data)
return result
@ -165,9 +165,9 @@ async def create_new_note(
request: Request,
form_data: NoteForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if user.role != 'admin' and not has_permission(
if user.role != 'admin' and not await has_permission(
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
):
raise HTTPException(
@ -175,7 +175,7 @@ async def create_new_note(
detail=ERROR_MESSAGES.UNAUTHORIZED,
)
form_data.access_grants = filter_allowed_access_grants(
form_data.access_grants = await filter_allowed_access_grants(
request.app.state.config.USER_PERMISSIONS,
user.id,
user.role,
@ -185,7 +185,7 @@ async def create_new_note(
)
try:
note = Notes.insert_new_note(user.id, form_data, db=db)
note = await Notes.insert_new_note(user.id, form_data, db=db)
return note
except Exception as e:
log.exception(e)
@ -206,9 +206,9 @@ async def get_note_by_id(
request: Request,
id: str,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if user.role != 'admin' and not has_permission(
if user.role != 'admin' and not await has_permission(
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
):
raise HTTPException(
@ -216,14 +216,14 @@ async def get_note_by_id(
detail=ERROR_MESSAGES.UNAUTHORIZED,
)
note = Notes.get_note_by_id(id, db=db)
note = await Notes.get_note_by_id(id, db=db)
if not note:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND)
if user.role != 'admin' and (
user.id != note.user_id
and (
not AccessGrants.has_access(
not await AccessGrants.has_access(
user_id=user.id,
resource_type='note',
resource_id=note.id,
@ -237,7 +237,7 @@ async def get_note_by_id(
write_access = (
user.role == 'admin'
or (user.id == note.user_id)
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user.id,
resource_type='note',
resource_id=note.id,
@ -261,9 +261,9 @@ async def update_note_by_id(
id: str,
form_data: NoteForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if user.role != 'admin' and not has_permission(
if user.role != 'admin' and not await has_permission(
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
):
raise HTTPException(
@ -271,13 +271,13 @@ async def update_note_by_id(
detail=ERROR_MESSAGES.UNAUTHORIZED,
)
note = Notes.get_note_by_id(id, db=db)
note = await Notes.get_note_by_id(id, db=db)
if not note:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND)
if user.role != 'admin' and (
user.id != note.user_id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='note',
resource_id=note.id,
@ -287,7 +287,7 @@ async def update_note_by_id(
):
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT())
form_data.access_grants = filter_allowed_access_grants(
form_data.access_grants = await filter_allowed_access_grants(
request.app.state.config.USER_PERMISSIONS,
user.id,
user.role,
@ -297,7 +297,7 @@ async def update_note_by_id(
)
try:
note = Notes.update_note_by_id(id, form_data, db=db)
note = await Notes.update_note_by_id(id, form_data, db=db)
await sio.emit(
'note-events',
note.model_dump(),
@ -325,9 +325,9 @@ async def update_note_access_by_id(
id: str,
form_data: NoteAccessGrantsForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if user.role != 'admin' and not has_permission(
if user.role != 'admin' and not await has_permission(
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
):
raise HTTPException(
@ -335,13 +335,13 @@ async def update_note_access_by_id(
detail=ERROR_MESSAGES.UNAUTHORIZED,
)
note = Notes.get_note_by_id(id, db=db)
note = await Notes.get_note_by_id(id, db=db)
if not note:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND)
if user.role != 'admin' and (
user.id != note.user_id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='note',
resource_id=note.id,
@ -351,7 +351,7 @@ async def update_note_access_by_id(
):
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT())
form_data.access_grants = filter_allowed_access_grants(
form_data.access_grants = await filter_allowed_access_grants(
request.app.state.config.USER_PERMISSIONS,
user.id,
user.role,
@ -359,9 +359,9 @@ async def update_note_access_by_id(
'sharing.public_notes',
)
AccessGrants.set_access_grants('note', id, form_data.access_grants, db=db)
await AccessGrants.set_access_grants('note', id, form_data.access_grants, db=db)
return Notes.get_note_by_id(id, db=db)
return await Notes.get_note_by_id(id, db=db)
############################
@ -374,9 +374,9 @@ async def delete_note_by_id(
request: Request,
id: str,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if user.role != 'admin' and not has_permission(
if user.role != 'admin' and not await has_permission(
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
):
raise HTTPException(
@ -384,13 +384,13 @@ async def delete_note_by_id(
detail=ERROR_MESSAGES.UNAUTHORIZED,
)
note = Notes.get_note_by_id(id, db=db)
note = await Notes.get_note_by_id(id, db=db)
if not note:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND)
if user.role != 'admin' and (
user.id != note.user_id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='note',
resource_id=note.id,
@ -401,7 +401,7 @@ async def delete_note_by_id(
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT())
try:
note = Notes.delete_note_by_id(id, db=db)
note = await Notes.delete_note_by_id(id, db=db)
return True
except Exception as e:
log.exception(e)

View file

@ -39,9 +39,9 @@ from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import StreamingResponse
from pydantic import BaseModel, ConfigDict, validator
from sqlalchemy.orm import Session
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import get_session
from open_webui.internal.db import get_async_session
from open_webui.models.models import Models
@ -398,11 +398,11 @@ async def get_all_models(request: Request, user: UserModel = None):
async def get_filtered_models(models, user, db=None):
# Filter models based on user access control
model_ids = [model['model'] for model in models.get('models', [])]
model_infos = {model_info.id: model_info for model_info in Models.get_models_by_ids(model_ids, db=db)}
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id, db=db)}
model_infos = {model_info.id: model_info for model_info in await Models.get_models_by_ids(model_ids, db=db)}
user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)}
# Batch-fetch accessible resource IDs in a single query instead of N has_access calls
accessible_model_ids = AccessGrants.get_accessible_resource_ids(
accessible_model_ids = await AccessGrants.get_accessible_resource_ids(
user_id=user.id,
resource_type='model',
resource_ids=list(model_infos.keys()),
@ -797,11 +797,14 @@ async def show_model_info(request: Request, form_data: ModelNameForm, user=Depen
form_data = form_data.model_dump(exclude_none=True)
form_data['model'] = form_data.get('model', form_data.get('name'))
model = form_data.get('model')
# Enforce per-model access control
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,
@ -846,6 +849,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
@ -902,6 +908,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
@ -964,11 +973,15 @@ async def generate_completion(
if not request.app.state.config.ENABLE_OLLAMA_API:
raise HTTPException(status_code=503, detail='Ollama API is disabled')
# Enforce per-model access control
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:
@ -1051,7 +1064,7 @@ async def generate_chat_completion(
if not request.app.state.config.ENABLE_OLLAMA_API:
raise HTTPException(status_code=503, detail='Ollama API is disabled')
# NOTE: We intentionally do NOT use Depends(get_session) here.
# NOTE: We intentionally do NOT use Depends(get_async_session) here.
# Database operations (get_model_by_id, AccessGrants.has_access) manage their own short-lived sessions.
# This prevents holding a connection during the entire LLM call (30-60+ seconds),
# which would exhaust the connection pool under concurrent load.
@ -1080,7 +1093,7 @@ async def generate_chat_completion(
del payload['metadata']
model_id = payload['model']
model_info = Models.get_model_by_id(model_id)
model_info = await Models.get_model_by_id(model_id)
if model_info:
if model_info.base_model_id:
@ -1098,9 +1111,9 @@ async def generate_chat_completion(
if not bypass_system_prompt:
payload = apply_system_prompt_to_body(system, payload, metadata, user)
check_model_access(user, model_info, bypass_filter)
await check_model_access(user, model_info, bypass_filter)
else:
check_model_access(user, None, bypass_filter)
await check_model_access(user, None, bypass_filter)
url, url_idx = await get_ollama_url(request, payload['model'], url_idx)
api_config = request.app.state.config.OLLAMA_API_CONFIGS.get(
@ -1158,7 +1171,7 @@ async def generate_openai_completion(
url_idx: Optional[int] = None,
user=Depends(get_verified_user),
):
# NOTE: We intentionally do NOT use Depends(get_session) here.
# NOTE: We intentionally do NOT use Depends(get_async_session) here.
# Database operations (get_model_by_id, AccessGrants.has_access) manage their own short-lived sessions.
# This prevents holding a connection during the entire LLM call (30-60+ seconds),
# which would exhaust the connection pool under concurrent load.
@ -1178,7 +1191,7 @@ async def generate_openai_completion(
del payload['metadata']
model_id = form_data.model
model_info = Models.get_model_by_id(model_id)
model_info = await Models.get_model_by_id(model_id)
if model_info:
if model_info.base_model_id:
payload['model'] = model_info.base_model_id
@ -1187,9 +1200,9 @@ async def generate_openai_completion(
if params:
payload = apply_model_params_to_body_openai(params, payload)
check_model_access(user, model_info)
await check_model_access(user, model_info)
else:
check_model_access(user, None)
await check_model_access(user, None)
url, url_idx = await get_ollama_url(request, payload['model'], url_idx)
api_config = request.app.state.config.OLLAMA_API_CONFIGS.get(
@ -1220,7 +1233,7 @@ async def generate_openai_chat_completion(
url_idx: Optional[int] = None,
user=Depends(get_verified_user),
):
# NOTE: We intentionally do NOT use Depends(get_session) here.
# NOTE: We intentionally do NOT use Depends(get_async_session) here.
# Database operations (get_model_by_id, AccessGrants.has_access) manage their own short-lived sessions.
# This prevents holding a connection during the entire LLM call (30-60+ seconds),
# which would exhaust the connection pool under concurrent load.
@ -1240,7 +1253,7 @@ async def generate_openai_chat_completion(
del payload['metadata']
model_id = completion_form.model
model_info = Models.get_model_by_id(model_id)
model_info = await Models.get_model_by_id(model_id)
if model_info:
if model_info.base_model_id:
payload['model'] = model_info.base_model_id
@ -1253,9 +1266,9 @@ async def generate_openai_chat_completion(
payload = apply_model_params_to_body_openai(params, payload)
payload = apply_system_prompt_to_body(system, payload, metadata, user)
check_model_access(user, model_info)
await check_model_access(user, model_info)
else:
check_model_access(user, None)
await check_model_access(user, None)
url, url_idx = await get_ollama_url(request, payload['model'], url_idx)
api_config = request.app.state.config.OLLAMA_API_CONFIGS.get(
@ -1300,14 +1313,14 @@ async def generate_anthropic_messages(
payload = {**form_data}
model_id = payload.get('model', '')
model_info = Models.get_model_by_id(model_id)
model_info = await Models.get_model_by_id(model_id)
if model_info:
if model_info.base_model_id:
payload['model'] = model_info.base_model_id
check_model_access(user, model_info)
await check_model_access(user, model_info)
else:
check_model_access(user, None)
await check_model_access(user, None)
url, url_idx = await get_ollama_url(request, payload['model'], url_idx)
api_config = request.app.state.config.OLLAMA_API_CONFIGS.get(
@ -1358,17 +1371,17 @@ async def generate_responses(
payload = form_data.model_dump()
model_id = form_data.model
model_info = Models.get_model_by_id(model_id)
model_info = await Models.get_model_by_id(model_id)
if model_info:
if model_info.base_model_id:
payload['model'] = model_info.base_model_id
# Check if user has access to the model
if user.role == 'user':
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)}
user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)}
if not (
user.id == model_info.user_id
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user.id,
resource_type='model',
resource_id=model_info.id,
@ -1413,7 +1426,7 @@ async def get_openai_models(
request: Request,
url_idx: Optional[int] = None,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
models = []
if url_idx is None:
@ -1445,11 +1458,11 @@ async def get_openai_models(
if user.role == 'user' and not BYPASS_MODEL_ACCESS_CONTROL:
# Filter models based on user access control
model_ids = [model['id'] for model in models]
model_infos = {model_info.id: model_info for model_info in Models.get_models_by_ids(model_ids, db=db)}
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id, db=db)}
model_infos = {model_info.id: model_info for model_info in await Models.get_models_by_ids(model_ids, db=db)}
user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)}
# Batch-fetch accessible resource IDs in a single query instead of N has_access calls
accessible_model_ids = AccessGrants.get_accessible_resource_ids(
accessible_model_ids = await AccessGrants.get_accessible_resource_ids(
user_id=user.id,
resource_type='model',
resource_ids=list(model_infos.keys()),
@ -1572,7 +1585,7 @@ async def download_model(
file_name = parse_huggingface_url(form_data.url)
if file_name:
file_path = f'{UPLOAD_DIR}/{file_name}'
file_path = os.path.join(UPLOAD_DIR, file_name)
return StreamingResponse(
download_file_stream(url, form_data.url, file_path, file_name),

View file

@ -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,9 +21,9 @@ from fastapi.responses import (
)
from pydantic import BaseModel, ConfigDict
from sqlalchemy.orm import Session
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import get_session
from open_webui.internal.db import get_async_session
from open_webui.models.models import Models
from open_webui.models.access_grants import AccessGrants
@ -40,6 +40,7 @@ from open_webui.env import (
ENABLE_FORWARD_USER_INFO_HEADERS,
FORWARD_SESSION_INFO_HEADER_CHAT_ID,
BYPASS_MODEL_ACCESS_CONTROL,
ENABLE_OPENAI_API_PASSTHROUGH,
)
from open_webui.models.users import UserModel
@ -450,11 +451,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()),
@ -772,6 +773,21 @@ def is_openai_new_model(model: str) -> bool:
return False
def _sanitize_model_for_url(model: str) -> str:
"""Sanitize a model name before interpolating it into a URL path.
Rejects path traversal attempts (../, /, \\) and percent-encodes
the name so it is safe to use as a single URL path segment
(e.g. Azure deployment name).
"""
if not model or '..' in model or '/' in model or '\\' in model:
raise HTTPException(
status_code=400,
detail='Invalid model name: must not be empty or contain path separators or traversal sequences',
)
return quote(model, safe='')
def convert_to_azure_payload(url, payload: dict, api_version: str):
model = payload.get('model', '')
@ -795,6 +811,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
@ -1007,7 +1026,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.
@ -1025,7 +1044,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:
@ -1045,9 +1064,9 @@ async def generate_chat_completion(
if not bypass_system_prompt:
payload = apply_system_prompt_to_body(system, payload, metadata, user)
check_model_access(user, model_info, bypass_filter)
await check_model_access(user, model_info, bypass_filter)
else:
check_model_access(user, None, bypass_filter)
await check_model_access(user, None, bypass_filter)
# Check if model is already in app state cache to avoid expensive get_all_models() call
models = request.app.state.OPENAI_MODELS
@ -1326,7 +1345,7 @@ async def responses(
model_id = form_data.model
# Enforce per-model access control
check_model_access(user, Models.get_model_by_id(model_id), BYPASS_MODEL_ACCESS_CONTROL)
await check_model_access(user, await Models.get_model_by_id(model_id), BYPASS_MODEL_ACCESS_CONTROL)
body = json.dumps(payload)
@ -1364,7 +1383,7 @@ async def responses(
else:
api_version = api_config.get('api_version', '2023-03-15-preview')
headers['api-version'] = api_version
model = payload.get('model', '')
model = _sanitize_model_for_url(payload.get('model', ''))
request_url = f'{url}/openai/deployments/{model}/responses?api-version={api_version}'
else:
request_url = f'{url}/responses'
@ -1404,6 +1423,8 @@ async def responses(
return response_data
except HTTPException:
raise
except Exception as e:
log.exception(e)
raise HTTPException(
@ -1418,9 +1439,16 @@ async def responses(
@router.api_route('/{path:path}', methods=['GET', 'POST', 'PUT', 'DELETE'])
async def proxy(path: str, request: Request, user=Depends(get_verified_user)):
"""
Deprecated: proxy all requests to OpenAI API
Deprecated: proxy all requests to OpenAI API.
Disabled by default. Set ENABLE_OPENAI_API_PASSTHROUGH=True to enable.
"""
if not ENABLE_OPENAI_API_PASSTHROUGH:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail='Direct API passthrough is disabled. Set ENABLE_OPENAI_API_PASSTHROUGH=True to enable.',
)
body = await request.body()
# Parse JSON body to resolve model-based routing
@ -1515,6 +1543,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(

View file

@ -20,8 +20,8 @@ from open_webui.constants import ERROR_MESSAGES
from open_webui.utils.auth import get_admin_user, get_verified_user
from open_webui.utils.access_control import has_permission, filter_allowed_access_grants
from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL
from open_webui.internal.db import get_session
from sqlalchemy.orm import Session
from open_webui.internal.db import get_async_session
from sqlalchemy.ext.asyncio import AsyncSession
from pydantic import BaseModel
@ -48,21 +48,21 @@ PAGE_ITEM_COUNT = 30
@router.get('/', response_model=list[PromptModel])
async def get_prompts(user=Depends(get_verified_user), db: Session = Depends(get_session)):
async def get_prompts(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
if user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL:
prompts = Prompts.get_prompts(db=db)
prompts = await Prompts.get_prompts(db=db)
else:
prompts = Prompts.get_prompts_by_user_id(user.id, 'read', db=db)
prompts = await Prompts.get_prompts_by_user_id(user.id, 'read', db=db)
return prompts
@router.get('/tags', response_model=list[str])
async def get_prompt_tags(user=Depends(get_verified_user), db: Session = Depends(get_session)):
async def get_prompt_tags(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
if user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL:
return Prompts.get_tags(db=db)
return await Prompts.get_tags(db=db)
else:
prompts = Prompts.get_prompts_by_user_id(user.id, 'read', db=db)
prompts = await Prompts.get_prompts_by_user_id(user.id, 'read', db=db)
tags = set()
for prompt in prompts:
if prompt.tags:
@ -79,7 +79,7 @@ async def get_prompt_list(
direction: Optional[str] = None,
page: Optional[int] = 1,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
limit = PAGE_ITEM_COUNT
@ -99,7 +99,7 @@ async def get_prompt_list(
filter['direction'] = direction
# Pre-fetch user group IDs once - used for both filter and write_access check
groups = Groups.get_groups_by_member_id(user.id, db=db)
groups = await Groups.get_groups_by_member_id(user.id, db=db)
user_group_ids = {group.id for group in groups}
if not (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL):
@ -108,11 +108,11 @@ async def get_prompt_list(
filter['user_id'] = user.id
result = Prompts.search_prompts(user.id, filter=filter, skip=skip, limit=limit, db=db)
result = await Prompts.search_prompts(user.id, filter=filter, skip=skip, limit=limit, db=db)
# Batch-fetch writable prompt IDs in a single query instead of N has_access calls
prompt_ids = [prompt.id for prompt in result.items]
writable_prompt_ids = AccessGrants.get_accessible_resource_ids(
writable_prompt_ids = await AccessGrants.get_accessible_resource_ids(
user_id=user.id,
resource_type='prompt',
resource_ids=prompt_ids,
@ -147,16 +147,16 @@ async def create_new_prompt(
request: Request,
form_data: PromptForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if user.role != 'admin' and not (
has_permission(
await has_permission(
user.id,
'workspace.prompts',
request.app.state.config.USER_PERMISSIONS,
db=db,
)
or has_permission(
or await has_permission(
user.id,
'workspace.prompts_import',
request.app.state.config.USER_PERMISSIONS,
@ -168,7 +168,7 @@ async def create_new_prompt(
detail=ERROR_MESSAGES.UNAUTHORIZED,
)
form_data.access_grants = filter_allowed_access_grants(
form_data.access_grants = await filter_allowed_access_grants(
request.app.state.config.USER_PERMISSIONS,
user.id,
user.role,
@ -176,9 +176,9 @@ async def create_new_prompt(
'sharing.public_prompts',
)
prompt = Prompts.get_prompt_by_command(form_data.command, db=db)
prompt = await Prompts.get_prompt_by_command(form_data.command, db=db)
if prompt is None:
prompt = Prompts.insert_new_prompt(user.id, form_data, db=db)
prompt = await Prompts.insert_new_prompt(user.id, form_data, db=db)
if prompt:
return prompt
@ -198,14 +198,14 @@ async def create_new_prompt(
@router.get('/command/{command}', response_model=Optional[PromptAccessResponse])
async def get_prompt_by_command(command: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
prompt = Prompts.get_prompt_by_command(command, db=db)
async def get_prompt_by_command(command: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
prompt = await Prompts.get_prompt_by_command(command, db=db)
if prompt:
if (
user.role == 'admin'
or prompt.user_id == user.id
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user.id,
resource_type='prompt',
resource_id=prompt.id,
@ -218,7 +218,7 @@ async def get_prompt_by_command(command: str, user=Depends(get_verified_user), d
write_access=(
(user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL)
or user.id == prompt.user_id
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user.id,
resource_type='prompt',
resource_id=prompt.id,
@ -240,14 +240,14 @@ async def get_prompt_by_command(command: str, user=Depends(get_verified_user), d
@router.get('/id/{prompt_id}', response_model=Optional[PromptAccessResponse])
async def get_prompt_by_id(prompt_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
async def get_prompt_by_id(prompt_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
prompt = await Prompts.get_prompt_by_id(prompt_id, db=db)
if prompt:
if (
user.role == 'admin'
or prompt.user_id == user.id
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user.id,
resource_type='prompt',
resource_id=prompt.id,
@ -260,7 +260,7 @@ async def get_prompt_by_id(prompt_id: str, user=Depends(get_verified_user), db:
write_access=(
(user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL)
or user.id == prompt.user_id
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user.id,
resource_type='prompt',
resource_id=prompt.id,
@ -287,9 +287,9 @@ async def update_prompt_by_id(
prompt_id: str,
form_data: PromptForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
prompt = await Prompts.get_prompt_by_id(prompt_id, db=db)
if not prompt:
raise HTTPException(
@ -300,7 +300,7 @@ async def update_prompt_by_id(
# Is the user the original creator, in a group with write access, or an admin
if (
prompt.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='prompt',
resource_id=prompt.id,
@ -316,14 +316,14 @@ async def update_prompt_by_id(
# Check for command collision if command is being changed
if form_data.command != prompt.command:
existing_prompt = Prompts.get_prompt_by_command(form_data.command, db=db)
existing_prompt = await Prompts.get_prompt_by_command(form_data.command, db=db)
if existing_prompt and existing_prompt.id != prompt.id:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"Command '/{form_data.command}' is already in use by another prompt",
)
form_data.access_grants = filter_allowed_access_grants(
form_data.access_grants = await filter_allowed_access_grants(
request.app.state.config.USER_PERMISSIONS,
user.id,
user.role,
@ -332,7 +332,7 @@ async def update_prompt_by_id(
)
# Use the ID from the found prompt
updated_prompt = Prompts.update_prompt_by_id(prompt.id, form_data, user.id, db=db)
updated_prompt = await Prompts.update_prompt_by_id(prompt.id, form_data, user.id, db=db)
if updated_prompt:
return updated_prompt
else:
@ -352,10 +352,10 @@ async def update_prompt_metadata(
prompt_id: str,
form_data: PromptMetadataForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""Update prompt name and command only (no history created)."""
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
prompt = await Prompts.get_prompt_by_id(prompt_id, db=db)
if not prompt:
raise HTTPException(
@ -365,7 +365,7 @@ async def update_prompt_metadata(
if (
prompt.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='prompt',
resource_id=prompt.id,
@ -381,14 +381,14 @@ async def update_prompt_metadata(
# Check for command collision if command is being changed
if form_data.command != prompt.command:
existing_prompt = Prompts.get_prompt_by_command(form_data.command, db=db)
existing_prompt = await Prompts.get_prompt_by_command(form_data.command, db=db)
if existing_prompt and existing_prompt.id != prompt.id:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"Command '/{form_data.command}' is already in use",
)
updated_prompt = Prompts.update_prompt_metadata(prompt.id, form_data.name, form_data.command, form_data.tags, db=db)
updated_prompt = await Prompts.update_prompt_metadata(prompt.id, form_data.name, form_data.command, form_data.tags, db=db)
if updated_prompt:
return updated_prompt
else:
@ -403,9 +403,9 @@ async def set_prompt_version(
prompt_id: str,
form_data: PromptVersionUpdateForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
prompt = await Prompts.get_prompt_by_id(prompt_id, db=db)
if not prompt:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
@ -414,7 +414,7 @@ async def set_prompt_version(
if (
prompt.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='prompt',
resource_id=prompt.id,
@ -428,7 +428,7 @@ async def set_prompt_version(
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
updated_prompt = Prompts.update_prompt_version(prompt.id, form_data.version_id, db=db)
updated_prompt = await Prompts.update_prompt_version(prompt.id, form_data.version_id, db=db)
if updated_prompt:
return updated_prompt
else:
@ -453,9 +453,9 @@ async def update_prompt_access_by_id(
prompt_id: str,
form_data: PromptAccessGrantsForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
prompt = await Prompts.get_prompt_by_id(prompt_id, db=db)
if not prompt:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
@ -464,7 +464,7 @@ async def update_prompt_access_by_id(
if (
prompt.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='prompt',
resource_id=prompt.id,
@ -478,7 +478,7 @@ async def update_prompt_access_by_id(
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
form_data.access_grants = filter_allowed_access_grants(
form_data.access_grants = await filter_allowed_access_grants(
request.app.state.config.USER_PERMISSIONS,
user.id,
user.role,
@ -486,9 +486,9 @@ async def update_prompt_access_by_id(
'sharing.public_prompts',
)
AccessGrants.set_access_grants('prompt', prompt_id, form_data.access_grants, db=db)
await AccessGrants.set_access_grants('prompt', prompt_id, form_data.access_grants, db=db)
return Prompts.get_prompt_by_id(prompt_id, db=db)
return await Prompts.get_prompt_by_id(prompt_id, db=db)
############################
@ -497,8 +497,8 @@ async def update_prompt_access_by_id(
@router.post('/id/{prompt_id}/toggle', response_model=Optional[PromptModel])
async def toggle_prompt_active(prompt_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
async def toggle_prompt_active(prompt_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
prompt = await Prompts.get_prompt_by_id(prompt_id, db=db)
if not prompt:
raise HTTPException(
@ -508,7 +508,7 @@ async def toggle_prompt_active(prompt_id: str, user=Depends(get_verified_user),
if (
prompt.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='prompt',
resource_id=prompt.id,
@ -522,7 +522,7 @@ async def toggle_prompt_active(prompt_id: str, user=Depends(get_verified_user),
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
result = Prompts.toggle_prompt_active(prompt.id, db=db)
result = await Prompts.toggle_prompt_active(prompt.id, db=db)
if result:
return result
raise HTTPException(
@ -537,8 +537,8 @@ async def toggle_prompt_active(prompt_id: str, user=Depends(get_verified_user),
@router.delete('/id/{prompt_id}/delete', response_model=bool)
async def delete_prompt_by_id(prompt_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
async def delete_prompt_by_id(prompt_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
prompt = await Prompts.get_prompt_by_id(prompt_id, db=db)
if not prompt:
raise HTTPException(
@ -548,7 +548,7 @@ async def delete_prompt_by_id(prompt_id: str, user=Depends(get_verified_user), d
if (
prompt.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='prompt',
resource_id=prompt.id,
@ -562,7 +562,7 @@ async def delete_prompt_by_id(prompt_id: str, user=Depends(get_verified_user), d
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
result = Prompts.delete_prompt_by_id(prompt.id, db=db)
result = await Prompts.delete_prompt_by_id(prompt.id, db=db)
return result
@ -576,12 +576,12 @@ async def get_prompt_history(
prompt_id: str,
page: int = 0,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""Get version history for a prompt."""
PAGE_SIZE = 20
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
prompt = await Prompts.get_prompt_by_id(prompt_id, db=db)
if not prompt:
raise HTTPException(
@ -593,7 +593,7 @@ async def get_prompt_history(
if not (
user.role == 'admin'
or prompt.user_id == user.id
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user.id,
resource_type='prompt',
resource_id=prompt.id,
@ -606,7 +606,7 @@ async def get_prompt_history(
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
history = PromptHistories.get_history_by_prompt_id(prompt.id, limit=PAGE_SIZE, offset=page * PAGE_SIZE, db=db)
history = await PromptHistories.get_history_by_prompt_id(prompt.id, limit=PAGE_SIZE, offset=page * PAGE_SIZE, db=db)
return history
@ -615,10 +615,10 @@ async def get_prompt_history_entry(
prompt_id: str,
history_id: str,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""Get a specific version from history."""
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
prompt = await Prompts.get_prompt_by_id(prompt_id, db=db)
if not prompt:
raise HTTPException(
@ -630,7 +630,7 @@ async def get_prompt_history_entry(
if not (
user.role == 'admin'
or prompt.user_id == user.id
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user.id,
resource_type='prompt',
resource_id=prompt.id,
@ -643,7 +643,7 @@ async def get_prompt_history_entry(
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
history_entry = PromptHistories.get_history_entry_by_id(history_id, db=db)
history_entry = await PromptHistories.get_history_entry_by_id(history_id, db=db)
if not history_entry or history_entry.prompt_id != prompt.id:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
@ -658,10 +658,10 @@ async def delete_prompt_history_entry(
prompt_id: str,
history_id: str,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""Delete a history entry. Cannot delete the active production version."""
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
prompt = await Prompts.get_prompt_by_id(prompt_id, db=db)
if not prompt:
raise HTTPException(
@ -673,7 +673,7 @@ async def delete_prompt_history_entry(
if not (
user.role == 'admin'
or prompt.user_id == user.id
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user.id,
resource_type='prompt',
resource_id=prompt.id,
@ -693,7 +693,7 @@ async def delete_prompt_history_entry(
detail='Cannot delete the active production version',
)
success = PromptHistories.delete_history_entry(history_id, db=db)
success = await PromptHistories.delete_history_entry(history_id, db=db)
if not success:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
@ -709,10 +709,10 @@ async def get_prompt_diff(
from_id: str,
to_id: str,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""Get diff between two versions."""
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
prompt = await Prompts.get_prompt_by_id(prompt_id, db=db)
if not prompt:
raise HTTPException(
@ -724,7 +724,7 @@ async def get_prompt_diff(
if not (
user.role == 'admin'
or prompt.user_id == user.id
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user.id,
resource_type='prompt',
resource_id=prompt.id,
@ -737,7 +737,7 @@ async def get_prompt_diff(
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
diff = PromptHistories.compute_diff(from_id, to_id, db=db)
diff = await PromptHistories.compute_diff(from_id, to_id, db=db)
if not diff:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,

View file

@ -41,8 +41,8 @@ from open_webui.utils.access_control.files import has_access_to_file
from open_webui.models.knowledge import Knowledges
from open_webui.models.access_grants import AccessGrants
from open_webui.storage.provider import Storage
from open_webui.internal.db import get_session, get_db
from sqlalchemy.orm import Session
from open_webui.internal.db import get_async_session, get_db
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.retrieval.vector.factory import VECTOR_DB_CLIENT
@ -1569,11 +1569,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.
@ -1582,9 +1582,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:
collection_name = form_data.collection_name
@ -1724,7 +1724,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,
@ -1732,8 +1732,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,
@ -1765,7 +1765,7 @@ def process_file(
if result:
# Fresh session for the final update.
with get_db() as session:
Files.update_file_metadata_by_id(
await Files.update_file_metadata_by_id(
file.id,
{
'collection_name': collection_name,
@ -1773,12 +1773,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,
@ -1795,13 +1795,13 @@ def process_file(
log.exception(e)
# Fresh session for error status update.
with get_db() as session:
Files.update_file_data_by_id(
await Files.update_file_data_by_id(
file.id,
{'status': 'failed'},
db=session,
)
# Clear the hash so the file can be re-uploaded after fixing the issue
Files.update_file_hash_by_id(file.id, None, db=session)
await Files.update_file_hash_by_id(file.id, None, db=session)
if 'No pandoc was found' in str(e):
raise HTTPException(
@ -2234,7 +2234,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(
@ -2386,7 +2386,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.
@ -2403,7 +2403,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,
@ -2429,7 +2429,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):
@ -2494,7 +2494,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):
@ -2555,14 +2555,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,
@ -2583,13 +2583,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
@ -2644,7 +2644,7 @@ async def process_files_batch(
"""
Process a batch of files and save them to the vector database.
NOTE: We intentionally do NOT use Depends(get_session) here.
NOTE: We intentionally do NOT use Depends(get_async_session) here.
The save_docs_to_vector_db() call makes external embedding API calls which
can take 5-60+ seconds for batch operations. Database operations after
embedding (Files.update_file_by_id) manage their own short-lived sessions.
@ -2662,7 +2662,7 @@ async def process_files_batch(
for file in form_data.files:
try:
# Ownership check: verify the requesting user owns the file or is an admin
db_file = Files.get_file_by_id(file.id, db=db)
db_file = await Files.get_file_by_id(file.id, db=db)
if not db_file:
file_errors.append(
BatchProcessFilesResult(
@ -2724,7 +2724,7 @@ async def process_files_batch(
# Update all files with collection name
for file_update, file_result in zip(file_updates, file_results):
Files.update_file_by_id(id=file_result.file_id, form_data=file_update, db=db)
await Files.update_file_by_id(id=file_result.file_id, form_data=file_update, db=db)
file_result.status = 'completed'
except Exception as e:

View file

@ -30,8 +30,8 @@ from open_webui.config import OAUTH_PROVIDERS
from open_webui.env import SCIM_AUTH_PROVIDER
from sqlalchemy.orm import Session
from open_webui.internal.db import get_session
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import get_async_session
log = logging.getLogger(__name__)
@ -326,18 +326,18 @@ def get_scim_provider() -> str:
return SCIM_AUTH_PROVIDER
def find_user_by_external_id(external_id: str, db=None) -> Optional[UserModel]:
async def find_user_by_external_id(external_id: str, db=None) -> Optional[UserModel]:
"""Find a user by SCIM externalId, falling back to OAuth sub match."""
provider = get_scim_provider()
user = Users.get_user_by_scim_external_id(provider, external_id, db=db)
user = await Users.get_user_by_scim_external_id(provider, external_id, db=db)
if user:
return user
# Fallback: check if externalId matches an existing OAuth sub (account linking)
return Users.get_user_by_oauth_sub(provider, external_id, db=db)
return await Users.get_user_by_oauth_sub(provider, external_id, db=db)
def user_to_scim(user: UserModel, request: Request, db=None) -> SCIMUser:
async def user_to_scim(user: UserModel, request: Request, db=None) -> SCIMUser:
"""Convert internal User model to SCIM User"""
# Parse display name into name components
name_parts = user.name.split(' ', 1) if user.name else ['', '']
@ -345,7 +345,7 @@ def user_to_scim(user: UserModel, request: Request, db=None) -> SCIMUser:
family_name = name_parts[1] if len(name_parts) > 1 else ''
# Get user's groups
user_groups = Groups.get_groups_by_member_id(user.id, db=db)
user_groups = await Groups.get_groups_by_member_id(user.id, db=db)
groups = [
{
'value': group.id,
@ -379,12 +379,12 @@ def user_to_scim(user: UserModel, request: Request, db=None) -> SCIMUser:
)
def group_to_scim(group: GroupModel, request: Request, db=None) -> SCIMGroup:
async def group_to_scim(group: GroupModel, request: Request, db=None) -> SCIMGroup:
"""Convert internal Group model to SCIM Group"""
member_ids = Groups.get_group_user_ids_by_id(group.id, db) or []
member_ids = await Groups.get_group_user_ids_by_id(group.id, db) or []
# Batch-fetch all users to avoid N+1 queries
users = Users.get_users_by_user_ids(member_ids, db=db) if member_ids else []
users = await Users.get_users_by_user_ids(member_ids, db=db) if member_ids else []
members = [
SCIMGroupMember(
value=user.id,
@ -512,7 +512,7 @@ async def get_users(
count: int = Query(20),
filter: Optional[str] = None,
_: bool = Depends(get_scim_auth),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""List SCIM Users"""
# Clamp per SCIM 2.0 spec (RFC 7644 §3.4.2.4):
@ -527,25 +527,25 @@ async def get_users(
# Simple filter parsing - supports userName eq, externalId eq
if 'userName eq' in filter:
email = filter.split('"')[1]
user = Users.get_user_by_email(email, db=db)
user = await Users.get_user_by_email(email, db=db)
users_list = [user] if user else []
total = 1 if user else 0
elif 'externalId eq' in filter:
external_id = filter.split('"')[1]
user = find_user_by_external_id(external_id, db=db)
user = await find_user_by_external_id(external_id, db=db)
users_list = [user] if user else []
total = 1 if user else 0
else:
response = Users.get_users(skip=skip, limit=limit, db=db)
response = await Users.get_users(skip=skip, limit=limit, db=db)
users_list = response['users']
total = response['total']
else:
response = Users.get_users(skip=skip, limit=limit, db=db)
response = await Users.get_users(skip=skip, limit=limit, db=db)
users_list = response['users']
total = response['total']
# Convert to SCIM format
scim_users = [user_to_scim(user, request, db=db) for user in users_list]
scim_users = [await user_to_scim(user, request, db=db) for user in users_list]
return SCIMListResponse(
totalResults=total,
@ -560,14 +560,14 @@ async def get_user(
user_id: str,
request: Request,
_: bool = Depends(get_scim_auth),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""Get SCIM User by ID"""
user = Users.get_user_by_id(user_id, db=db)
user = await Users.get_user_by_id(user_id, db=db)
if not user:
return scim_error(status_code=status.HTTP_404_NOT_FOUND, detail=f'User {user_id} not found')
return user_to_scim(user, request, db=db)
return await user_to_scim(user, request, db=db)
@router.post('/Users', response_model=SCIMUser, status_code=status.HTTP_201_CREATED)
@ -575,12 +575,12 @@ async def create_user(
request: Request,
user_data: SCIMUserCreateRequest,
_: bool = Depends(get_scim_auth),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""Create SCIM User"""
# Check for duplicate by externalId
if user_data.externalId:
existing_user = find_user_by_external_id(user_data.externalId, db=db)
existing_user = await find_user_by_external_id(user_data.externalId, db=db)
if existing_user:
raise HTTPException(
status_code=status.HTTP_409_CONFLICT,
@ -596,7 +596,7 @@ async def create_user(
email = email.lower()
# Check for duplicate by email
existing_user = Users.get_user_by_email(email, db=db)
existing_user = await Users.get_user_by_email(email, db=db)
if existing_user:
raise HTTPException(
status_code=status.HTTP_409_CONFLICT,
@ -619,7 +619,7 @@ async def create_user(
if user_data.photos and len(user_data.photos) > 0:
profile_image = user_data.photos[0].value
new_user = Users.insert_new_user(
new_user = await Users.insert_new_user(
id=user_id,
name=name,
email=email,
@ -637,10 +637,10 @@ async def create_user(
# Store externalId in the scim field
if user_data.externalId:
provider = get_scim_provider()
Users.update_user_scim_by_id(user_id, provider, user_data.externalId, db=db)
new_user = Users.get_user_by_id(user_id, db=db)
await Users.update_user_scim_by_id(user_id, provider, user_data.externalId, db=db)
new_user = await Users.get_user_by_id(user_id, db=db)
return user_to_scim(new_user, request, db=db)
return await user_to_scim(new_user, request, db=db)
@router.put('/Users/{user_id}', response_model=SCIMUser)
@ -649,10 +649,10 @@ async def update_user(
request: Request,
user_data: SCIMUserUpdateRequest,
_: bool = Depends(get_scim_auth),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""Update SCIM User (full update)"""
user = Users.get_user_by_id(user_id, db=db)
user = await Users.get_user_by_id(user_id, db=db)
if not user:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
@ -682,7 +682,7 @@ async def update_user(
if user_data.photos and len(user_data.photos) > 0:
update_data['profile_image_url'] = user_data.photos[0].value
updated_user = Users.update_user_by_id(user_id, update_data, db=db)
updated_user = await Users.update_user_by_id(user_id, update_data, db=db)
if not updated_user:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
@ -692,10 +692,10 @@ async def update_user(
# Update externalId in the scim field
if user_data.externalId:
provider = get_scim_provider()
Users.update_user_scim_by_id(user_id, provider, user_data.externalId, db=db)
updated_user = Users.get_user_by_id(user_id, db=db)
await Users.update_user_scim_by_id(user_id, provider, user_data.externalId, db=db)
updated_user = await Users.get_user_by_id(user_id, db=db)
return user_to_scim(updated_user, request, db=db)
return await user_to_scim(updated_user, request, db=db)
@router.patch('/Users/{user_id}', response_model=SCIMUser)
@ -704,10 +704,10 @@ async def patch_user(
request: Request,
patch_data: SCIMPatchRequest,
_: bool = Depends(get_scim_auth),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""Update SCIM User (partial update)"""
user = Users.get_user_by_id(user_id, db=db)
user = await Users.get_user_by_id(user_id, db=db)
if not user:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
@ -734,11 +734,11 @@ async def patch_user(
update_data['name'] = value
elif path == 'externalId':
provider = get_scim_provider()
Users.update_user_scim_by_id(user_id, provider, value, db=db)
await Users.update_user_scim_by_id(user_id, provider, value, db=db)
# Update user
if update_data:
updated_user = Users.update_user_by_id(user_id, update_data, db=db)
updated_user = await Users.update_user_by_id(user_id, update_data, db=db)
if not updated_user:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
@ -747,7 +747,7 @@ async def patch_user(
else:
updated_user = user
return user_to_scim(updated_user, request, db=db)
return await user_to_scim(updated_user, request, db=db)
@router.delete('/Users/{user_id}', status_code=status.HTTP_204_NO_CONTENT)
@ -755,17 +755,17 @@ async def delete_user(
user_id: str,
request: Request,
_: bool = Depends(get_scim_auth),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""Delete SCIM User"""
user = Users.get_user_by_id(user_id, db=db)
user = await Users.get_user_by_id(user_id, db=db)
if not user:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=f'User {user_id} not found',
)
success = Users.delete_user_by_id(user_id, db=db)
success = await Users.delete_user_by_id(user_id, db=db)
if not success:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
@ -783,7 +783,7 @@ async def get_groups(
count: int = Query(20),
filter: Optional[str] = None,
_: bool = Depends(get_scim_auth),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""List SCIM Groups"""
# Clamp per SCIM 2.0 spec (RFC 7644 §3.4.2.4):
@ -795,13 +795,13 @@ async def get_groups(
if filter:
if 'displayName eq' in filter:
display_name = filter.split('"')[1]
group = Groups.get_group_by_name(display_name, db=db)
group = await Groups.get_group_by_name(display_name, db=db)
groups_list = [group] if group else []
else:
# Unrecognized filter — fall back to all groups
groups_list = Groups.get_all_groups(db=db)
groups_list = await Groups.get_all_groups(db=db)
else:
groups_list = Groups.get_all_groups(db=db)
groups_list = await Groups.get_all_groups(db=db)
# Apply pagination
total = len(groups_list)
@ -810,7 +810,7 @@ async def get_groups(
paginated_groups = groups_list[start:end]
# Convert to SCIM format
scim_groups = [group_to_scim(group, request, db=db) for group in paginated_groups]
scim_groups = [await group_to_scim(group, request, db=db) for group in paginated_groups]
return SCIMListResponse(
totalResults=total,
@ -825,17 +825,17 @@ async def get_group(
group_id: str,
request: Request,
_: bool = Depends(get_scim_auth),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""Get SCIM Group by ID"""
group = Groups.get_group_by_id(group_id, db=db)
group = await Groups.get_group_by_id(group_id, db=db)
if not group:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=f'Group {group_id} not found',
)
return group_to_scim(group, request, db=db)
return await group_to_scim(group, request, db=db)
@router.post('/Groups', response_model=SCIMGroup, status_code=status.HTTP_201_CREATED)
@ -843,7 +843,7 @@ async def create_group(
request: Request,
group_data: SCIMGroupCreateRequest,
_: bool = Depends(get_scim_auth),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""Create SCIM Group"""
# Extract member IDs
@ -861,14 +861,14 @@ async def create_group(
)
# Need to get the creating user's ID - we'll use the first admin
admin_user = Users.get_super_admin_user(db=db)
admin_user = await Users.get_super_admin_user(db=db)
if not admin_user:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail='No admin user found',
)
new_group = Groups.insert_new_group(admin_user.id, form, db=db)
new_group = await Groups.insert_new_group(admin_user.id, form, db=db)
if not new_group:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
@ -884,12 +884,12 @@ async def create_group(
description=new_group.description,
)
Groups.update_group_by_id(new_group.id, update_form, db=db)
Groups.set_group_user_ids_by_id(new_group.id, member_ids, db=db)
await Groups.update_group_by_id(new_group.id, update_form, db=db)
await Groups.set_group_user_ids_by_id(new_group.id, member_ids, db=db)
new_group = Groups.get_group_by_id(new_group.id, db=db)
new_group = await Groups.get_group_by_id(new_group.id, db=db)
return group_to_scim(new_group, request, db=db)
return await group_to_scim(new_group, request, db=db)
@router.put('/Groups/{group_id}', response_model=SCIMGroup)
@ -898,10 +898,10 @@ async def update_group(
request: Request,
group_data: SCIMGroupUpdateRequest,
_: bool = Depends(get_scim_auth),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""Update SCIM Group (full update)"""
group = Groups.get_group_by_id(group_id, db=db)
group = await Groups.get_group_by_id(group_id, db=db)
if not group:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
@ -919,17 +919,17 @@ async def update_group(
# Handle members if provided
if group_data.members is not None:
member_ids = [member.value for member in group_data.members]
Groups.set_group_user_ids_by_id(group_id, member_ids, db=db)
await Groups.set_group_user_ids_by_id(group_id, member_ids, db=db)
# Update group
updated_group = Groups.update_group_by_id(group_id, update_form, db=db)
updated_group = await Groups.update_group_by_id(group_id, update_form, db=db)
if not updated_group:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail='Failed to update group',
)
return group_to_scim(updated_group, request, db=db)
return await group_to_scim(updated_group, request, db=db)
@router.patch('/Groups/{group_id}', response_model=SCIMGroup)
@ -938,10 +938,10 @@ async def patch_group(
request: Request,
patch_data: SCIMPatchRequest,
_: bool = Depends(get_scim_auth),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""Update SCIM Group (partial update)"""
group = Groups.get_group_by_id(group_id, db=db)
group = await Groups.get_group_by_id(group_id, db=db)
if not group:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
@ -965,7 +965,7 @@ async def patch_group(
update_form.name = value
elif path == 'members':
# Replace all members
Groups.set_group_user_ids_by_id(group_id, [member['value'] for member in value], db=db)
await Groups.set_group_user_ids_by_id(group_id, [member['value'] for member in value], db=db)
elif op == 'add':
if path == 'members':
@ -973,22 +973,22 @@ async def patch_group(
if isinstance(value, list):
for member in value:
if isinstance(member, dict) and 'value' in member:
Groups.add_users_to_group(group_id, [member['value']], db=db)
await Groups.add_users_to_group(group_id, [member['value']], db=db)
elif op == 'remove':
if path and path.startswith('members[value eq'):
# Remove specific member
member_id = path.split('"')[1]
Groups.remove_users_from_group(group_id, [member_id], db=db)
await Groups.remove_users_from_group(group_id, [member_id], db=db)
# Update group
updated_group = Groups.update_group_by_id(group_id, update_form, db=db)
updated_group = await Groups.update_group_by_id(group_id, update_form, db=db)
if not updated_group:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail='Failed to update group',
)
return group_to_scim(updated_group, request, db=db)
return await group_to_scim(updated_group, request, db=db)
@router.delete('/Groups/{group_id}', status_code=status.HTTP_204_NO_CONTENT)
@ -996,17 +996,17 @@ async def delete_group(
group_id: str,
request: Request,
_: bool = Depends(get_scim_auth),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
"""Delete SCIM Group"""
group = Groups.get_group_by_id(group_id, db=db)
group = await Groups.get_group_by_id(group_id, db=db)
if not group:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=f'Group {group_id} not found',
)
success = Groups.delete_group_by_id(group_id, db=db)
success = await Groups.delete_group_by_id(group_id, db=db)
if not success:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,

View file

@ -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

View file

@ -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

View file

@ -8,8 +8,8 @@ from open_webui.env import AIOHTTP_CLIENT_TIMEOUT
from open_webui.models.groups import Groups
from pydantic import BaseModel, HttpUrl
from fastapi import APIRouter, Depends, HTTPException, Request, status
from sqlalchemy.orm import Session
from open_webui.internal.db import get_session
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import get_async_session
from open_webui.models.oauth_sessions import OAuthSessions
@ -46,11 +46,11 @@ log = logging.getLogger(__name__)
router = APIRouter()
def get_tool_module(request, tool_id, load_from_db=True):
async def get_tool_module(request, tool_id, load_from_db=True):
"""
Get the tool module by its ID.
"""
tool_module, _ = get_tool_module_from_cache(request, tool_id, load_from_db)
tool_module, _ = await get_tool_module_from_cache(request, tool_id, load_from_db)
return tool_module
@ -65,12 +65,12 @@ def get_tool_module(request, tool_id, load_from_db=True):
async def get_tools(
request: Request,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
tools = []
# Local Tools
for tool in Tools.get_tools(defer_content=True, db=db):
for tool in await Tools.get_tools(defer_content=True, db=db):
tool_module = request.app.state.TOOLS.get(tool.id) if hasattr(request.app.state, 'TOOLS') else None
tools.append(
ToolUserResponse(
@ -159,31 +159,30 @@ async def get_tools(
# Admin can see all tools
return tools
else:
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id, db=db)}
tools = [
tool
for tool in tools
if tool.user_id == user.id
or (
has_access(
user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)}
filtered_tools = []
for tool in tools:
if tool.user_id == user.id:
filtered_tools.append(tool)
elif str(tool.id).startswith('server:'):
if await has_access(
user.id,
'read',
server_access_grants.get(str(tool.id), []),
user_group_ids,
db=db,
)
if str(tool.id).startswith('server:')
else AccessGrants.has_access(
user_id=user.id,
resource_type='tool',
resource_id=tool.id,
permission='read',
user_group_ids=user_group_ids,
db=db,
)
)
]
return tools
):
filtered_tools.append(tool)
elif await AccessGrants.has_access(
user_id=user.id,
resource_type='tool',
resource_id=tool.id,
permission='read',
user_group_ids=user_group_ids,
db=db,
):
filtered_tools.append(tool)
return filtered_tools
############################
@ -192,13 +191,13 @@ async def get_tools(
@router.get('/list', response_model=list[ToolAccessResponse])
async def get_tool_list(user=Depends(get_verified_user), db: Session = Depends(get_session)):
async def get_tool_list(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
if user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL:
tools = Tools.get_tools(defer_content=True, db=db)
tools = await Tools.get_tools(defer_content=True, db=db)
else:
tools = Tools.get_tools_by_user_id(user.id, 'read', defer_content=True, db=db)
tools = await Tools.get_tools_by_user_id(user.id, 'read', defer_content=True, db=db)
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id, db=db)}
user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)}
result = []
for tool in tools:
@ -298,9 +297,9 @@ async def load_tool_from_url(request: Request, form_data: LoadUrlForm, user=Depe
async def export_tools(
request: Request,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if user.role != 'admin' and not has_permission(
if user.role != 'admin' and not await has_permission(
user.id,
'workspace.tools_export',
request.app.state.config.USER_PERMISSIONS,
@ -312,9 +311,9 @@ async def export_tools(
)
if user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL:
return Tools.get_tools(db=db)
return await Tools.get_tools(db=db)
else:
return Tools.get_tools_by_user_id(user.id, 'read', db=db)
return await Tools.get_tools_by_user_id(user.id, 'read', db=db)
############################
@ -327,11 +326,11 @@ async def create_new_tools(
request: Request,
form_data: ToolForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if user.role != 'admin' and not (
has_permission(user.id, 'workspace.tools', request.app.state.config.USER_PERMISSIONS, db=db)
or has_permission(
await has_permission(user.id, 'workspace.tools', request.app.state.config.USER_PERMISSIONS, db=db)
or await has_permission(
user.id,
'workspace.tools_import',
request.app.state.config.USER_PERMISSIONS,
@ -351,10 +350,10 @@ async def create_new_tools(
form_data.id = form_data.id.lower()
tools = Tools.get_tool_by_id(form_data.id, db=db)
tools = await Tools.get_tool_by_id(form_data.id, db=db)
if tools is None:
try:
form_data.access_grants = filter_allowed_access_grants(
form_data.access_grants = await filter_allowed_access_grants(
request.app.state.config.USER_PERMISSIONS,
user.id,
user.role,
@ -363,14 +362,14 @@ async def create_new_tools(
)
form_data.content = replace_imports(form_data.content)
tool_module, frontmatter = load_tool_module_by_id(form_data.id, content=form_data.content)
tool_module, frontmatter = await load_tool_module_by_id(form_data.id, content=form_data.content)
form_data.meta.manifest = frontmatter
TOOLS = request.app.state.TOOLS
TOOLS[form_data.id] = tool_module
specs = get_tool_specs(TOOLS[form_data.id])
tools = Tools.insert_new_tool(user.id, form_data, specs, db=db)
tools = await Tools.insert_new_tool(user.id, form_data, specs, db=db)
tool_cache_dir = CACHE_DIR / 'tools' / form_data.id
tool_cache_dir.mkdir(parents=True, exist_ok=True)
@ -401,14 +400,14 @@ async def create_new_tools(
@router.get('/id/{id}', response_model=Optional[ToolAccessResponse])
async def get_tools_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
tools = Tools.get_tool_by_id(id, db=db)
async def get_tools_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
tools = await Tools.get_tool_by_id(id, db=db)
if tools:
if (
user.role == 'admin'
or tools.user_id == user.id
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user.id,
resource_type='tool',
resource_id=tools.id,
@ -421,7 +420,7 @@ async def get_tools_by_id(id: str, user=Depends(get_verified_user), db: Session
write_access=(
(user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL)
or user.id == tools.user_id
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user.id,
resource_type='tool',
resource_id=tools.id,
@ -453,9 +452,9 @@ async def update_tools_by_id(
id: str,
form_data: ToolForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
tools = Tools.get_tool_by_id(id, db=db)
tools = await Tools.get_tool_by_id(id, db=db)
if not tools:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
@ -465,7 +464,7 @@ async def update_tools_by_id(
# Is the user the original creator, in a group with write access, or an admin
if (
tools.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='tool',
resource_id=tools.id,
@ -481,7 +480,7 @@ async def update_tools_by_id(
try:
form_data.content = replace_imports(form_data.content)
tool_module, frontmatter = load_tool_module_by_id(id, content=form_data.content)
tool_module, frontmatter = await load_tool_module_by_id(id, content=form_data.content)
form_data.meta.manifest = frontmatter
TOOLS = request.app.state.TOOLS
@ -489,7 +488,7 @@ async def update_tools_by_id(
specs = get_tool_specs(TOOLS[id])
form_data.access_grants = filter_allowed_access_grants(
form_data.access_grants = await filter_allowed_access_grants(
request.app.state.config.USER_PERMISSIONS,
user.id,
user.role,
@ -503,7 +502,7 @@ async def update_tools_by_id(
}
log.debug(updated)
tools = Tools.update_tool_by_id(id, updated, db=db)
tools = await Tools.update_tool_by_id(id, updated, db=db)
if tools:
return tools
@ -535,9 +534,9 @@ async def update_tool_access_by_id(
id: str,
form_data: ToolAccessGrantsForm,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
tools = Tools.get_tool_by_id(id, db=db)
tools = await Tools.get_tool_by_id(id, db=db)
if not tools:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
@ -546,7 +545,7 @@ async def update_tool_access_by_id(
if (
tools.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='tool',
resource_id=tools.id,
@ -560,7 +559,7 @@ async def update_tool_access_by_id(
detail=ERROR_MESSAGES.UNAUTHORIZED,
)
form_data.access_grants = filter_allowed_access_grants(
form_data.access_grants = await filter_allowed_access_grants(
request.app.state.config.USER_PERMISSIONS,
user.id,
user.role,
@ -568,9 +567,9 @@ async def update_tool_access_by_id(
'sharing.public_tools',
)
AccessGrants.set_access_grants('tool', id, form_data.access_grants, db=db)
await AccessGrants.set_access_grants('tool', id, form_data.access_grants, db=db)
return Tools.get_tool_by_id(id, db=db)
return await Tools.get_tool_by_id(id, db=db)
############################
@ -583,9 +582,9 @@ async def delete_tools_by_id(
request: Request,
id: str,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
tools = Tools.get_tool_by_id(id, db=db)
tools = await Tools.get_tool_by_id(id, db=db)
if not tools:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
@ -594,7 +593,7 @@ async def delete_tools_by_id(
if (
tools.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='tool',
resource_id=tools.id,
@ -608,7 +607,7 @@ async def delete_tools_by_id(
detail=ERROR_MESSAGES.UNAUTHORIZED,
)
result = Tools.delete_tool_by_id(id, db=db)
result = await Tools.delete_tool_by_id(id, db=db)
if result:
TOOLS = request.app.state.TOOLS
if id in TOOLS:
@ -623,8 +622,8 @@ async def delete_tools_by_id(
@router.get('/id/{id}/valves', response_model=Optional[dict])
async def get_tools_valves_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
tools = Tools.get_tool_by_id(id, db=db)
async def get_tools_valves_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
tools = await Tools.get_tool_by_id(id, db=db)
if not tools:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
@ -633,7 +632,7 @@ async def get_tools_valves_by_id(id: str, user=Depends(get_verified_user), db: S
if (
tools.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='tool',
resource_id=tools.id,
@ -648,7 +647,7 @@ async def get_tools_valves_by_id(id: str, user=Depends(get_verified_user), db: S
)
try:
valves = Tools.get_tool_valves_by_id(id, db=db)
valves = await Tools.get_tool_valves_by_id(id, db=db)
return valves
except Exception as e:
raise HTTPException(
@ -667,9 +666,9 @@ async def get_tools_valves_spec_by_id(
request: Request,
id: str,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
tools = Tools.get_tool_by_id(id, db=db)
tools = await Tools.get_tool_by_id(id, db=db)
if not tools:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
@ -678,7 +677,7 @@ async def get_tools_valves_spec_by_id(
if (
tools.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='tool',
resource_id=tools.id,
@ -695,7 +694,7 @@ async def get_tools_valves_spec_by_id(
if id in request.app.state.TOOLS:
tools_module = request.app.state.TOOLS[id]
else:
tools_module, _ = load_tool_module_by_id(id)
tools_module, _ = await load_tool_module_by_id(id)
request.app.state.TOOLS[id] = tools_module
if hasattr(tools_module, 'Valves'):
@ -718,9 +717,9 @@ async def update_tools_valves_by_id(
id: str,
form_data: dict,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
tools = Tools.get_tool_by_id(id, db=db)
tools = await Tools.get_tool_by_id(id, db=db)
if not tools:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
@ -729,7 +728,7 @@ async def update_tools_valves_by_id(
if (
tools.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='tool',
resource_id=tools.id,
@ -746,7 +745,7 @@ async def update_tools_valves_by_id(
if id in request.app.state.TOOLS:
tools_module = request.app.state.TOOLS[id]
else:
tools_module, _ = load_tool_module_by_id(id)
tools_module, _ = await load_tool_module_by_id(id)
request.app.state.TOOLS[id] = tools_module
if not hasattr(tools_module, 'Valves'):
@ -760,7 +759,7 @@ async def update_tools_valves_by_id(
form_data = {k: v for k, v in form_data.items() if v is not None}
valves = Valves(**form_data)
valves_dict = valves.model_dump(exclude_unset=True)
Tools.update_tool_valves_by_id(id, valves_dict, db=db)
await Tools.update_tool_valves_by_id(id, valves_dict, db=db)
return valves_dict
except Exception as e:
log.exception(f'Failed to update tool valves by id {id}: {e}')
@ -776,8 +775,8 @@ async def update_tools_valves_by_id(
@router.get('/id/{id}/valves/user', response_model=Optional[dict])
async def get_tools_user_valves_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
tools = Tools.get_tool_by_id(id, db=db)
async def get_tools_user_valves_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
tools = await Tools.get_tool_by_id(id, db=db)
if not tools:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
@ -786,7 +785,7 @@ async def get_tools_user_valves_by_id(id: str, user=Depends(get_verified_user),
if (
tools.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='tool',
resource_id=tools.id,
@ -801,7 +800,7 @@ async def get_tools_user_valves_by_id(id: str, user=Depends(get_verified_user),
)
try:
user_valves = Tools.get_user_valves_by_id_and_user_id(id, user.id, db=db)
user_valves = await Tools.get_user_valves_by_id_and_user_id(id, user.id, db=db)
return user_valves
except Exception as e:
raise HTTPException(
@ -815,9 +814,9 @@ async def get_tools_user_valves_spec_by_id(
request: Request,
id: str,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
tools = Tools.get_tool_by_id(id, db=db)
tools = await Tools.get_tool_by_id(id, db=db)
if not tools:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
@ -826,7 +825,7 @@ async def get_tools_user_valves_spec_by_id(
if (
tools.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='tool',
resource_id=tools.id,
@ -843,7 +842,7 @@ async def get_tools_user_valves_spec_by_id(
if id in request.app.state.TOOLS:
tools_module = request.app.state.TOOLS[id]
else:
tools_module, _ = load_tool_module_by_id(id)
tools_module, _ = await load_tool_module_by_id(id)
request.app.state.TOOLS[id] = tools_module
if hasattr(tools_module, 'UserValves'):
@ -861,9 +860,9 @@ async def update_tools_user_valves_by_id(
id: str,
form_data: dict,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
tools = Tools.get_tool_by_id(id, db=db)
tools = await Tools.get_tool_by_id(id, db=db)
if not tools:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
@ -872,7 +871,7 @@ async def update_tools_user_valves_by_id(
if (
tools.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='tool',
resource_id=tools.id,
@ -889,7 +888,7 @@ async def update_tools_user_valves_by_id(
if id in request.app.state.TOOLS:
tools_module = request.app.state.TOOLS[id]
else:
tools_module, _ = load_tool_module_by_id(id)
tools_module, _ = await load_tool_module_by_id(id)
request.app.state.TOOLS[id] = tools_module
if hasattr(tools_module, 'UserValves'):
@ -899,7 +898,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}')

View file

@ -1,6 +1,6 @@
import logging
from typing import Optional
from sqlalchemy.orm import Session
from sqlalchemy.ext.asyncio import AsyncSession
import base64
import io
@ -30,7 +30,7 @@ from open_webui.models.users import (
from open_webui.constants import ERROR_MESSAGES
from open_webui.env import STATIC_DIR
from open_webui.internal.db import get_session
from open_webui.internal.db import get_async_session
from open_webui.utils.auth import (
@ -63,7 +63,7 @@ async def get_users(
direction: Optional[str] = None,
page: Optional[int] = 1,
user=Depends(get_admin_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
limit = PAGE_ITEM_COUNT
@ -80,14 +80,14 @@ async def get_users(
filter['direction'] = direction
result = Users.get_users(filter=filter, skip=skip, limit=limit, db=db)
result = await Users.get_users(filter=filter, skip=skip, limit=limit, db=db)
users = result['users']
total = result['total']
# Fetch groups for all users in a single query to avoid N+1
user_ids = [user.id for user in users]
user_groups = Groups.get_groups_by_member_ids(user_ids, db=db)
user_groups = await Groups.get_groups_by_member_ids(user_ids, db=db)
return {
'users': [
@ -106,9 +106,9 @@ async def get_users(
@router.get('/all', response_model=UserInfoListResponse)
async def get_all_users(
user=Depends(get_admin_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
return Users.get_users(db=db)
return await Users.get_users(db=db)
@router.get('/search', response_model=UserInfoListResponse)
@ -118,7 +118,7 @@ async def search_users(
direction: Optional[str] = None,
page: Optional[int] = 1,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
limit = PAGE_ITEM_COUNT
@ -133,7 +133,7 @@ async def search_users(
if direction:
filter['direction'] = direction
return Users.get_users(filter=filter, skip=skip, limit=limit, db=db)
return await Users.get_users(filter=filter, skip=skip, limit=limit, db=db)
############################
@ -142,8 +142,8 @@ async def search_users(
@router.get('/groups')
async def get_user_groups(user=Depends(get_verified_user), db: Session = Depends(get_session)):
return Groups.get_groups_by_member_id(user.id, db=db)
async def get_user_groups(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
return await Groups.get_groups_by_member_id(user.id, db=db)
############################
@ -155,9 +155,9 @@ async def get_user_groups(user=Depends(get_verified_user), db: Session = Depends
async def get_user_permissisions(
request: Request,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
user_permissions = get_permissions(user.id, request.app.state.config.USER_PERMISSIONS, db=db)
user_permissions = await get_permissions(user.id, request.app.state.config.USER_PERMISSIONS, db=db)
return user_permissions
@ -272,8 +272,8 @@ async def update_default_user_permissions(request: Request, form_data: UserPermi
@router.get('/user/settings', response_model=Optional[UserSettings])
async def get_user_settings_by_session_user(user=Depends(get_verified_user), db: Session = Depends(get_session)):
user = Users.get_user_by_id(user.id, db=db)
async def get_user_settings_by_session_user(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
user = await Users.get_user_by_id(user.id, db=db)
if user:
return user.settings
else:
@ -293,7 +293,7 @@ async def update_user_settings_by_session_user(
request: Request,
form_data: UserSettings,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
updated_user_settings = form_data.model_dump()
ui_settings = updated_user_settings.get('ui')
@ -301,7 +301,7 @@ async def update_user_settings_by_session_user(
user.role != 'admin'
and ui_settings is not None
and 'toolServers' in ui_settings.keys()
and not has_permission(
and not await has_permission(
user.id,
'features.direct_tool_servers',
request.app.state.config.USER_PERMISSIONS,
@ -310,7 +310,7 @@ async def update_user_settings_by_session_user(
# If the user is not an admin and does not have permission to use tool servers, remove the key
updated_user_settings['ui'].pop('toolServers', None)
user = Users.update_user_settings_by_id(user.id, updated_user_settings, db=db)
user = await Users.update_user_settings_by_id(user.id, updated_user_settings, db=db)
if user:
return user.settings
else:
@ -329,14 +329,14 @@ async def update_user_settings_by_session_user(
async def get_user_status_by_session_user(
request: Request,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if not request.app.state.config.ENABLE_USER_STATUS:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=ERROR_MESSAGES.ACTION_PROHIBITED,
)
user = Users.get_user_by_id(user.id, db=db)
user = await Users.get_user_by_id(user.id, db=db)
if user:
return user
else:
@ -356,16 +356,16 @@ async def update_user_status_by_session_user(
request: Request,
form_data: UserStatus,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
if not request.app.state.config.ENABLE_USER_STATUS:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=ERROR_MESSAGES.ACTION_PROHIBITED,
)
user = Users.get_user_by_id(user.id, db=db)
user = await Users.get_user_by_id(user.id, db=db)
if user:
user = Users.update_user_status_by_id(user.id, form_data, db=db)
user = await Users.update_user_status_by_id(user.id, form_data, db=db)
return user
else:
raise HTTPException(
@ -380,8 +380,8 @@ async def update_user_status_by_session_user(
@router.get('/user/info', response_model=Optional[dict])
async def get_user_info_by_session_user(user=Depends(get_verified_user), db: Session = Depends(get_session)):
user = Users.get_user_by_id(user.id, db=db)
async def get_user_info_by_session_user(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
user = await Users.get_user_by_id(user.id, db=db)
if user:
return user.info
else:
@ -398,14 +398,14 @@ async def get_user_info_by_session_user(user=Depends(get_verified_user), db: Ses
@router.post('/user/info/update', response_model=Optional[dict])
async def update_user_info_by_session_user(
form_data: dict, user=Depends(get_verified_user), db: Session = Depends(get_session)
form_data: dict, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
user = Users.get_user_by_id(user.id, db=db)
user = await Users.get_user_by_id(user.id, db=db)
if user:
if user.info is None:
user.info = {}
user = Users.update_user_by_id(user.id, {'info': {**user.info, **form_data}}, db=db)
user = await Users.update_user_by_id(user.id, {'info': {**user.info, **form_data}}, db=db)
if user:
return user.info
else:
@ -435,12 +435,12 @@ class UserActiveResponse(UserStatus):
@router.get('/{user_id}', response_model=UserActiveResponse)
async def get_user_by_id(user_id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)):
async def get_user_by_id(user_id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
# Check if user_id is a shared chat
# If it is, get the user_id from the chat
if user_id.startswith('shared-'):
chat_id = user_id.replace('shared-', '')
chat = Chats.get_chat_by_id(chat_id)
chat = await Chats.get_chat_by_id(chat_id)
if chat:
user_id = chat.user_id
else:
@ -449,14 +449,14 @@ async def get_user_by_id(user_id: str, user=Depends(get_admin_user), db: Session
detail=ERROR_MESSAGES.USER_NOT_FOUND,
)
user = Users.get_user_by_id(user_id, db=db)
user = await Users.get_user_by_id(user_id, db=db)
if user:
groups = Groups.get_groups_by_member_id(user_id, db=db)
groups = await Groups.get_groups_by_member_id(user_id, db=db)
return UserActiveResponse(
**{
**user.model_dump(),
'groups': [{'id': group.id, 'name': group.name} for group in groups],
'is_active': Users.is_user_active(user_id, db=db),
'is_active': await Users.is_user_active(user_id, db=db),
}
)
else:
@ -467,15 +467,15 @@ async def get_user_by_id(user_id: str, user=Depends(get_admin_user), db: Session
@router.get('/{user_id}/info', response_model=UserInfoResponse)
async def get_user_info_by_id(user_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
user = Users.get_user_by_id(user_id, db=db)
async def get_user_info_by_id(user_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
user = await Users.get_user_by_id(user_id, db=db)
if user:
groups = Groups.get_groups_by_member_id(user_id, db=db)
groups = await Groups.get_groups_by_member_id(user_id, db=db)
return UserInfoResponse(
**{
**user.model_dump(),
'groups': [{'id': group.id, 'name': group.name} for group in groups],
'is_active': Users.is_user_active(user_id, db=db),
'is_active': await Users.is_user_active(user_id, db=db),
}
)
else:
@ -486,8 +486,8 @@ async def get_user_info_by_id(user_id: str, user=Depends(get_verified_user), db:
@router.get('/{user_id}/oauth/sessions')
async def get_user_oauth_sessions_by_id(user_id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)):
sessions = OAuthSessions.get_sessions_by_user_id(user_id, db=db)
async def get_user_oauth_sessions_by_id(user_id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
sessions = await OAuthSessions.get_sessions_by_user_id(user_id, db=db)
if sessions and len(sessions) > 0:
return sessions
else:
@ -503,8 +503,8 @@ async def get_user_oauth_sessions_by_id(user_id: str, user=Depends(get_admin_use
@router.get('/{user_id}/profile/image')
def get_user_profile_image_by_id(user_id: str, user=Depends(get_verified_user)):
user = Users.get_user_by_id(user_id)
async def get_user_profile_image_by_id(user_id: str, user=Depends(get_verified_user)):
user = await Users.get_user_by_id(user_id)
if user:
if user.profile_image_url:
# check if it's url or base64
@ -542,10 +542,10 @@ def get_user_profile_image_by_id(user_id: str, user=Depends(get_verified_user)):
@router.get('/{user_id}/active', response_model=dict)
async def get_user_active_status_by_id(
user_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
user_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
return {
'active': Users.is_user_active(user_id, db=db),
'active': await Users.is_user_active(user_id, db=db),
}
@ -559,11 +559,11 @@ async def update_user_by_id(
user_id: str,
form_data: UserUpdateForm,
session_user=Depends(get_admin_user),
db: Session = Depends(get_session),
db: AsyncSession = Depends(get_async_session),
):
# Prevent modification of the primary admin user by other admins
try:
first_user = Users.get_first_user(db=db)
first_user = await Users.get_first_user(db=db)
if first_user:
if user_id == first_user.id:
if session_user.id != user_id:
@ -587,11 +587,11 @@ async def update_user_by_id(
detail='Could not verify primary admin status.',
)
user = Users.get_user_by_id(user_id, db=db)
user = await Users.get_user_by_id(user_id, db=db)
if user:
if form_data.email.lower() != user.email:
email_user = Users.get_user_by_email(form_data.email.lower(), db=db)
email_user = await Users.get_user_by_email(form_data.email.lower(), db=db)
if email_user:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
@ -605,10 +605,10 @@ async def update_user_by_id(
raise HTTPException(400, detail=str(e))
hashed = get_password_hash(form_data.password)
Auths.update_user_password_by_id(user_id, hashed, db=db)
await Auths.update_user_password_by_id(user_id, hashed, db=db)
Auths.update_email_by_id(user_id, form_data.email.lower(), db=db)
updated_user = Users.update_user_by_id(
await Auths.update_email_by_id(user_id, form_data.email.lower(), db=db)
updated_user = await Users.update_user_by_id(
user_id,
{
'role': form_data.role,
@ -639,10 +639,10 @@ async def update_user_by_id(
@router.delete('/{user_id}', response_model=bool)
async def delete_user_by_id(user_id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)):
async def delete_user_by_id(user_id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
# Prevent deletion of the primary admin user
try:
first_user = Users.get_first_user(db=db)
first_user = await Users.get_first_user(db=db)
if first_user and user_id == first_user.id:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
@ -656,7 +656,7 @@ async def delete_user_by_id(user_id: str, user=Depends(get_admin_user), db: Sess
)
if user.id != user_id:
result = Auths.delete_auth_by_id(user_id, db=db)
result = await Auths.delete_auth_by_id(user_id, db=db)
if result:
return True
@ -679,5 +679,5 @@ async def delete_user_by_id(user_id: str, user=Depends(get_admin_user), db: Sess
@router.get('/{user_id}/groups')
async def get_user_groups_by_id(user_id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)):
return Groups.get_groups_by_member_id(user_id, db=db)
async def get_user_groups_by_id(user_id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
return await Groups.get_groups_by_member_id(user_id, db=db)

View file

@ -333,7 +333,7 @@ async def connect(sid, environ, auth):
data = decode_token(auth['token'])
if data is not None and 'id' in data:
user = Users.get_user_by_id(data['id'])
user = await Users.get_user_by_id(data['id'])
if user:
SESSION_POOL[sid] = {
@ -361,7 +361,7 @@ async def user_join(sid, data):
if data is None or 'id' not in data:
return
user = Users.get_user_by_id(data['id'])
user = await Users.get_user_by_id(data['id'])
if not user:
return
@ -381,8 +381,8 @@ async def user_join(sid, data):
await sio.enter_room(sid, f'user:{user.id}')
# Join all the channels only if user has channels permission
if user.role == 'admin' or has_permission(user.id, 'features.channels'):
channels = Channels.get_channels_by_user_id(user.id)
if user.role == 'admin' or await has_permission(user.id, 'features.channels'):
channels = await Channels.get_channels_by_user_id(user.id)
log.debug(f'{channels=}')
for channel in channels:
await sio.enter_room(sid, f'channel:{channel.id}')
@ -395,7 +395,7 @@ async def heartbeat(sid, data):
user = SESSION_POOL.get(sid)
if user:
SESSION_POOL[sid] = {**user, 'last_seen_at': int(time.time())}
await asyncio.to_thread(Users.update_last_active_by_id, user['id'])
await Users.update_last_active_by_id(user['id'])
@sio.on('join-channels')
@ -408,13 +408,13 @@ async def join_channel(sid, data):
if data is None or 'id' not in data:
return
user = Users.get_user_by_id(data['id'])
user = await Users.get_user_by_id(data['id'])
if not user:
return
# Join all the channels only if user has channels permission
if user.role == 'admin' or has_permission(user.id, 'features.channels'):
channels = Channels.get_channels_by_user_id(user.id)
if user.role == 'admin' or await has_permission(user.id, 'features.channels'):
channels = await Channels.get_channels_by_user_id(user.id)
log.debug(f'{channels=}')
for channel in channels:
await sio.enter_room(sid, f'channel:{channel.id}')
@ -430,11 +430,11 @@ async def join_note(sid, data):
if token_data is None or 'id' not in token_data:
return
user = Users.get_user_by_id(token_data['id'])
user = await Users.get_user_by_id(token_data['id'])
if not user:
return
note = Notes.get_note_by_id(data['note_id'])
note = await Notes.get_note_by_id(data['note_id'])
if not note:
log.error(f'Note {data["note_id"]} not found for user {user.id}')
return
@ -442,7 +442,7 @@ async def join_note(sid, data):
if (
user.role != 'admin'
and user.id != note.user_id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='note',
resource_id=note.id,
@ -488,7 +488,7 @@ async def channel_events(sid, data):
room=room,
)
elif event_type == 'last_read_at':
Channels.update_member_last_read_at(data['channel_id'], user['id'])
await Channels.update_member_last_read_at(data['channel_id'], user['id'])
@sio.on('events:chat')
@ -501,7 +501,7 @@ async def chat_events(sid, data):
event_type = event_data.get('type')
if event_type == 'last_read_at':
await asyncio.to_thread(Chats.update_chat_last_read_at_by_id, data['chat_id'], user['id'])
await Chats.update_chat_last_read_at_by_id(data['chat_id'], user['id'])
def normalize_document_id(document_id: str) -> str:
@ -529,7 +529,7 @@ async def ydoc_document_join(sid, data):
if document_id.startswith('note:'):
note_id = document_id.split(':')[1]
note = Notes.get_note_by_id(note_id)
note = await Notes.get_note_by_id(note_id)
if not note:
log.error(f'Note {note_id} not found')
return
@ -537,7 +537,7 @@ async def ydoc_document_join(sid, data):
if (
user.get('role') != 'admin'
and user.get('id') != note.user_id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.get('id'),
resource_type='note',
resource_id=note.id,
@ -602,7 +602,7 @@ async def document_save_handler(document_id, data, user):
if document_id.startswith('note:'):
note_id = document_id.split(':')[1]
note = Notes.get_note_by_id(note_id)
note = await Notes.get_note_by_id(note_id)
if not note:
log.error(f'Note {note_id} not found')
return
@ -610,7 +610,7 @@ async def document_save_handler(document_id, data, user):
if (
user.get('role') != 'admin'
and user.get('id') != note.user_id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.get('id'),
resource_type='note',
resource_id=note.id,
@ -620,7 +620,7 @@ async def document_save_handler(document_id, data, user):
log.error(f'User {user.get("id")} does not have write access to note {note_id}')
return
Notes.update_note_by_id(note_id, NoteUpdateForm(data=data))
await Notes.update_note_by_id(note_id, NoteUpdateForm(data=data))
@sio.on('ydoc:document:state')
@ -793,7 +793,7 @@ async def disconnect(sid):
# print(f"Unknown session ID {sid} disconnected")
def get_event_emitter(request_info, update_db=True):
async def get_event_emitter(request_info, update_db=True):
async def __event_emitter__(event_data):
user_id = request_info['user_id']
chat_id = request_info['chat_id']
@ -813,16 +813,14 @@ def get_event_emitter(request_info, update_db=True):
event_type = event_data.get('type')
if event_type == 'status':
await asyncio.to_thread(
Chats.add_message_status_to_chat_by_id_and_message_id,
await Chats.add_message_status_to_chat_by_id_and_message_id(
request_info['chat_id'],
request_info['message_id'],
event_data.get('data', {}),
)
elif event_type == 'message':
message = await asyncio.to_thread(
Chats.get_message_by_id_and_message_id,
message = await Chats.get_message_by_id_and_message_id(
request_info['chat_id'],
request_info['message_id'],
)
@ -831,8 +829,7 @@ def get_event_emitter(request_info, update_db=True):
content = message.get('content', '')
content += event_data.get('data', {}).get('content', '')
await asyncio.to_thread(
Chats.upsert_message_to_chat_by_id_and_message_id,
await Chats.upsert_message_to_chat_by_id_and_message_id(
request_info['chat_id'],
request_info['message_id'],
{
@ -843,8 +840,7 @@ def get_event_emitter(request_info, update_db=True):
elif event_type == 'replace':
content = event_data.get('data', {}).get('content', '')
await asyncio.to_thread(
Chats.upsert_message_to_chat_by_id_and_message_id,
await Chats.upsert_message_to_chat_by_id_and_message_id(
request_info['chat_id'],
request_info['message_id'],
{
@ -853,8 +849,7 @@ def get_event_emitter(request_info, update_db=True):
)
elif event_type == 'embeds':
message = await asyncio.to_thread(
Chats.get_message_by_id_and_message_id,
message = await Chats.get_message_by_id_and_message_id(
request_info['chat_id'],
request_info['message_id'],
)
@ -862,8 +857,7 @@ def get_event_emitter(request_info, update_db=True):
embeds = event_data.get('data', {}).get('embeds', [])
embeds.extend(message.get('embeds', []))
await asyncio.to_thread(
Chats.upsert_message_to_chat_by_id_and_message_id,
await Chats.upsert_message_to_chat_by_id_and_message_id(
request_info['chat_id'],
request_info['message_id'],
{
@ -872,8 +866,7 @@ def get_event_emitter(request_info, update_db=True):
)
elif event_type == 'files':
message = await asyncio.to_thread(
Chats.get_message_by_id_and_message_id,
message = await Chats.get_message_by_id_and_message_id(
request_info['chat_id'],
request_info['message_id'],
)
@ -881,8 +874,7 @@ def get_event_emitter(request_info, update_db=True):
files = event_data.get('data', {}).get('files', [])
files.extend(message.get('files', []))
await asyncio.to_thread(
Chats.upsert_message_to_chat_by_id_and_message_id,
await Chats.upsert_message_to_chat_by_id_and_message_id(
request_info['chat_id'],
request_info['message_id'],
{
@ -893,8 +885,7 @@ def get_event_emitter(request_info, update_db=True):
elif event_type in ('source', 'citation'):
data = event_data.get('data', {})
if data.get('type') is None:
message = await asyncio.to_thread(
Chats.get_message_by_id_and_message_id,
message = await Chats.get_message_by_id_and_message_id(
request_info['chat_id'],
request_info['message_id'],
)
@ -902,8 +893,7 @@ def get_event_emitter(request_info, update_db=True):
sources = message.get('sources', [])
sources.append(data)
await asyncio.to_thread(
Chats.upsert_message_to_chat_by_id_and_message_id,
await Chats.upsert_message_to_chat_by_id_and_message_id(
request_info['chat_id'],
request_info['message_id'],
{
@ -917,7 +907,7 @@ def get_event_emitter(request_info, update_db=True):
return None
def get_event_call(request_info):
async def get_event_call(request_info):
async def __event_caller__(event_data):
response = await sio.call(
'events',

View file

@ -61,7 +61,7 @@ class LocalStorageProvider(StorageProvider):
contents = file.read()
if not contents:
raise ValueError(ERROR_MESSAGES.EMPTY_CONTENT)
file_path = f'{UPLOAD_DIR}/{filename}'
file_path = os.path.join(UPLOAD_DIR, filename)
with open(file_path, 'wb') as f:
f.write(contents)
return contents, file_path
@ -74,8 +74,8 @@ class LocalStorageProvider(StorageProvider):
@staticmethod
def delete_file(file_path: str) -> None:
"""Handles deletion of the file from local storage."""
filename = file_path.split('/')[-1]
file_path = f'{UPLOAD_DIR}/{filename}'
filename = os.path.basename(file_path)
file_path = os.path.join(UPLOAD_DIR, filename)
if os.path.isfile(file_path):
os.remove(file_path)
else:
@ -202,7 +202,7 @@ class S3StorageProvider(StorageProvider):
return '/'.join(full_file_path.split('//')[1].split('/')[1:])
def _get_local_file_path(self, s3_key: str) -> str:
return f'{UPLOAD_DIR}/{s3_key.split("/")[-1]}'
return os.path.join(UPLOAD_DIR, s3_key.split('/')[-1])
class GCSStorageProvider(StorageProvider):
@ -234,7 +234,7 @@ class GCSStorageProvider(StorageProvider):
"""Handles downloading of the file from GCS storage."""
try:
filename = file_path.removeprefix('gs://').split('/')[1]
local_file_path = f'{UPLOAD_DIR}/{filename}'
local_file_path = os.path.join(UPLOAD_DIR, filename)
blob = self.bucket.get_blob(filename)
blob.download_to_filename(local_file_path)
@ -298,7 +298,7 @@ class AzureStorageProvider(StorageProvider):
"""Handles downloading of the file from Azure Blob Storage."""
try:
filename = file_path.split('/')[-1]
local_file_path = f'{UPLOAD_DIR}/{filename}'
local_file_path = os.path.join(UPLOAD_DIR, filename)
blob_client = self.container_client.get_blob_client(filename)
with open(local_file_path, 'wb') as download_file:
download_file.write(blob_client.download_blob().readall())

View file

@ -250,7 +250,7 @@ async def generate_image(
# Persist files to DB if chat context is available
if __chat_id__ and __message_id__ and images:
db_files = Chats.add_message_files_by_id_and_message_id(
db_files = await Chats.add_message_files_by_id_and_message_id(
__chat_id__,
__message_id__,
image_files,
@ -317,7 +317,7 @@ async def edit_image(
# Persist files to DB if chat context is available
if __chat_id__ and __message_id__ and images:
db_files = Chats.add_message_files_by_id_and_message_id(
db_files = await Chats.add_message_files_by_id_and_message_id(
__chat_id__,
__message_id__,
image_files,
@ -473,14 +473,14 @@ async def execute_code(
from open_webui.models.users import Users
from open_webui.utils.files import get_image_url_from_base64
user = Users.get_user_by_id(__user__['id'])
user = await Users.get_user_by_id(__user__['id'])
# Extract and upload images from stdout
if stdout and isinstance(stdout, str):
stdout_lines = stdout.split('\n')
for idx, line in enumerate(stdout_lines):
if 'data:image/png;base64' in line:
image_url = get_image_url_from_base64(
image_url = await get_image_url_from_base64(
__request__,
line,
__metadata__ or {},
@ -495,7 +495,7 @@ async def execute_code(
result_lines = result.split('\n')
for idx, line in enumerate(result_lines):
if 'data:image/png;base64' in line:
image_url = get_image_url_from_base64(
image_url = await get_image_url_from_base64(
__request__,
line,
__metadata__ or {},
@ -650,7 +650,7 @@ async def delete_memory(
try:
user = UserModel(**__user__) if __user__ else None
result = Memories.delete_memory_by_id_and_user_id(memory_id, user.id)
result = await Memories.delete_memory_by_id_and_user_id(memory_id, user.id)
if result:
VECTOR_DB_CLIENT.delete(collection_name=f'user-memory-{user.id}', ids=[memory_id])
@ -680,7 +680,7 @@ async def list_memories(
try:
user = UserModel(**__user__) if __user__ else None
memories = Memories.get_memories_by_user_id(user.id)
memories = await Memories.get_memories_by_user_id(user.id)
if memories:
result = [
@ -730,9 +730,9 @@ async def search_notes(
try:
user_id = __user__.get('id')
user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id)]
user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)]
result = Notes.search_notes(
result = await Notes.search_notes(
user_id=user_id,
filter={
'query': query,
@ -808,18 +808,18 @@ async def view_note(
return json.dumps({'error': 'User context not available'})
try:
note = Notes.get_note_by_id(note_id)
note = await Notes.get_note_by_id(note_id)
if not note:
return json.dumps({'error': 'Note not found'})
# Check access permission
user_id = __user__.get('id')
user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id)]
user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)]
from open_webui.models.access_grants import AccessGrants
if note.user_id != user_id and not AccessGrants.has_access(
if note.user_id != user_id and not await AccessGrants.has_access(
user_id=user_id,
resource_type='note',
resource_id=note.id,
@ -878,7 +878,7 @@ async def write_note(
access_grants=[], # Private by default - only owner can access
)
new_note = Notes.insert_new_note(user_id, form)
new_note = await Notes.insert_new_note(user_id, form)
if not new_note:
return json.dumps({'error': 'Failed to create note'})
@ -921,18 +921,18 @@ async def replace_note_content(
try:
from open_webui.models.notes import NoteUpdateForm
note = Notes.get_note_by_id(note_id)
note = await Notes.get_note_by_id(note_id)
if not note:
return json.dumps({'error': 'Note not found'})
# Check write permission
user_id = __user__.get('id')
user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id)]
user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)]
from open_webui.models.access_grants import AccessGrants
if note.user_id != user_id and not AccessGrants.has_access(
if note.user_id != user_id and not await AccessGrants.has_access(
user_id=user_id,
resource_type='note',
resource_id=note.id,
@ -947,7 +947,7 @@ async def replace_note_content(
update_data['title'] = title
form = NoteUpdateForm(**update_data)
updated_note = Notes.update_note_by_id(note_id, form)
updated_note = await Notes.update_note_by_id(note_id, form)
if not updated_note:
return json.dumps({'error': 'Failed to update note'})
@ -998,7 +998,7 @@ async def search_chats(
try:
user_id = __user__.get('id')
chats = Chats.get_chats_by_user_id_and_search_text(
chats = await Chats.get_chats_by_user_id_and_search_text(
user_id=user_id,
search_text=query,
include_archived=False,
@ -1073,7 +1073,7 @@ async def view_chat(
try:
user_id = __user__.get('id')
chat = Chats.get_chat_by_id_and_user_id(chat_id, user_id)
chat = await Chats.get_chat_by_id_and_user_id(chat_id, user_id)
if not chat:
return json.dumps({'error': 'Chat not found or access denied'})
@ -1145,7 +1145,7 @@ async def search_channels(
user_id = __user__.get('id')
# Get all channels the user has access to
all_channels = Channels.get_channels_by_user_id(user_id)
all_channels = await Channels.get_channels_by_user_id(user_id)
# Filter by query
lower_query = query.lower()
@ -1201,7 +1201,7 @@ async def search_channel_messages(
user_id = __user__.get('id')
# Get all channels the user has access to
user_channels = Channels.get_channels_by_user_id(user_id)
user_channels = await Channels.get_channels_by_user_id(user_id)
channel_ids = [c.id for c in user_channels]
channel_map = {c.id: c for c in user_channels}
@ -1280,12 +1280,12 @@ async def view_channel_message(
return json.dumps({'error': 'Message not found'})
# Verify user has access to the channel
channel = Channels.get_channel_by_id(message.channel_id)
channel = await Channels.get_channel_by_id(message.channel_id)
if not channel:
return json.dumps({'error': 'Channel not found'})
# Check if user has access to the channel
user_channels = Channels.get_channels_by_user_id(user_id)
user_channels = await Channels.get_channels_by_user_id(user_id)
channel_ids = [c.id for c in user_channels]
if message.channel_id not in channel_ids:
@ -1342,11 +1342,11 @@ async def view_channel_thread(
return json.dumps({'error': 'Message not found'})
# Verify user has access to the channel
channel = Channels.get_channel_by_id(parent_message.channel_id)
channel = await Channels.get_channel_by_id(parent_message.channel_id)
if not channel:
return json.dumps({'error': 'Channel not found'})
user_channels = Channels.get_channels_by_user_id(user_id)
user_channels = await Channels.get_channels_by_user_id(user_id)
channel_ids = [c.id for c in user_channels]
if parent_message.channel_id not in channel_ids:
@ -1427,9 +1427,9 @@ async def list_knowledge_bases(
from open_webui.models.knowledge import Knowledges
user_id = __user__.get('id')
user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id)]
user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)]
result = Knowledges.search_knowledge_bases(
result = await Knowledges.search_knowledge_bases(
user_id,
filter={
'query': '',
@ -1442,7 +1442,7 @@ async def list_knowledge_bases(
knowledge_bases = []
for knowledge_base in result.items:
files = Knowledges.get_files_by_id(knowledge_base.id)
files = await Knowledges.get_files_by_id(knowledge_base.id)
file_count = len(files) if files else 0
knowledge_bases.append(
@ -1486,9 +1486,9 @@ async def search_knowledge_bases(
from open_webui.models.knowledge import Knowledges
user_id = __user__.get('id')
user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id)]
user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)]
result = Knowledges.search_knowledge_bases(
result = await Knowledges.search_knowledge_bases(
user_id,
filter={
'query': query,
@ -1501,7 +1501,7 @@ async def search_knowledge_bases(
knowledge_bases = []
for knowledge_base in result.items:
files = Knowledges.get_files_by_id(knowledge_base.id)
files = await Knowledges.get_files_by_id(knowledge_base.id)
file_count = len(files) if files else 0
knowledge_bases.append(
@ -1552,7 +1552,7 @@ async def search_knowledge_files(
user_id = __user__.get('id')
user_role = __user__.get('role', 'user')
user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id)]
user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)]
# When model has attached knowledge, scope to attached KBs/files only
if __model_knowledge__:
@ -1577,14 +1577,14 @@ async def search_knowledge_files(
# Search within attached KBs
for kb_id in attached_kb_ids:
knowledge = Knowledges.get_knowledge_by_id(kb_id)
knowledge = await Knowledges.get_knowledge_by_id(kb_id)
if not knowledge:
continue
if not (
user_role == 'admin'
or knowledge.user_id == user_id
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user_id,
resource_type='knowledge',
resource_id=knowledge.id,
@ -1594,7 +1594,7 @@ async def search_knowledge_files(
):
continue
result = Knowledges.search_files_by_id(
result = await Knowledges.search_files_by_id(
knowledge_id=kb_id,
user_id=user_id,
filter={'query': query},
@ -1617,7 +1617,7 @@ async def search_knowledge_files(
if not knowledge_id and attached_file_ids:
query_lower = query.lower() if query else ''
for file_id in attached_file_ids:
file = Files.get_file_by_id(file_id)
file = await Files.get_file_by_id(file_id)
if file and (not query_lower or query_lower in file.filename.lower()):
all_files.append(
{
@ -1633,7 +1633,7 @@ async def search_knowledge_files(
# No attached knowledge - search all accessible KBs
if knowledge_id:
result = Knowledges.search_files_by_id(
result = await Knowledges.search_files_by_id(
knowledge_id=knowledge_id,
user_id=user_id,
filter={'query': query},
@ -1641,7 +1641,7 @@ async def search_knowledge_files(
limit=count,
)
else:
result = Knowledges.search_knowledge_files(
result = await Knowledges.search_knowledge_files(
filter={
'query': query,
'user_id': user_id,
@ -1719,7 +1719,7 @@ async def view_file(
user_id = __user__.get('id')
user_role = __user__.get('role', 'user')
file = Files.get_file_by_id(file_id)
file = await Files.get_file_by_id(file_id)
if not file:
return json.dumps({'error': 'File not found'})
@ -1729,7 +1729,7 @@ async def view_file(
and not any(
item.get('type') == 'file' and item.get('id') == file_id for item in (__model_knowledge__ or [])
)
and not has_access_to_file(
and not await has_access_to_file(
file_id=file_id,
access_type='read',
user=UserModel(**__user__),
@ -1811,14 +1811,14 @@ async def view_knowledge_file(
user_id = __user__.get('id')
user_role = __user__.get('role', 'user')
user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id)]
user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)]
file = Files.get_file_by_id(file_id)
file = await Files.get_file_by_id(file_id)
if not file:
return json.dumps({'error': 'File not found'})
# Check access via any KB containing this file
knowledges = Knowledges.get_knowledges_by_file_id(file_id)
knowledges = await Knowledges.get_knowledges_by_file_id(file_id)
has_knowledge_access = False
knowledge_info = None
@ -1826,7 +1826,7 @@ async def view_knowledge_file(
if (
user_role == 'admin'
or knowledge_base.user_id == user_id
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user_id,
resource_type='knowledge',
resource_id=knowledge_base.id,
@ -1903,7 +1903,7 @@ async def list_knowledge(
user_id = __user__.get('id')
user_role = __user__.get('role', 'user')
user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id)]
user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)]
knowledge_bases = []
files = []
@ -1914,11 +1914,11 @@ async def list_knowledge(
item_id = item.get('id')
if item_type == 'collection':
knowledge = Knowledges.get_knowledge_by_id(item_id)
knowledge = await Knowledges.get_knowledge_by_id(item_id)
if knowledge and (
user_role == 'admin'
or knowledge.user_id == user_id
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user_id,
resource_type='knowledge',
resource_id=knowledge.id,
@ -1926,7 +1926,7 @@ async def list_knowledge(
user_group_ids=set(user_group_ids),
)
):
kb_files = Knowledges.get_files_by_id(knowledge.id)
kb_files = await Knowledges.get_files_by_id(knowledge.id)
file_count = len(kb_files) if kb_files else 0
kb_entry = {
@ -1943,7 +1943,7 @@ async def list_knowledge(
knowledge_bases.append(kb_entry)
elif item_type == 'file':
file = Files.get_file_by_id(item_id)
file = await Files.get_file_by_id(item_id)
if file:
files.append(
{
@ -1954,11 +1954,11 @@ async def list_knowledge(
)
elif item_type == 'note':
note = Notes.get_note_by_id(item_id)
note = await Notes.get_note_by_id(item_id)
if note and (
user_role == 'admin'
or note.user_id == user_id
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user_id,
resource_type='note',
resource_id=note.id,
@ -2036,7 +2036,7 @@ async def query_knowledge_files(
user_id = __user__.get('id')
user_role = __user__.get('role', 'user')
user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id)]
user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)]
embedding_function = __request__.app.state.EMBEDDING_FUNCTION
if not embedding_function:
@ -2053,11 +2053,11 @@ async def query_knowledge_files(
if item_type == 'collection':
# Knowledge base - use KB ID as collection name
knowledge = Knowledges.get_knowledge_by_id(item_id)
knowledge = await Knowledges.get_knowledge_by_id(item_id)
if knowledge and (
user_role == 'admin'
or knowledge.user_id == user_id
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user_id,
resource_type='knowledge',
resource_id=knowledge.id,
@ -2069,17 +2069,17 @@ async def query_knowledge_files(
elif item_type == 'file':
# Individual file - use file-{id} as collection name
file = Files.get_file_by_id(item_id)
file = await Files.get_file_by_id(item_id)
if file:
collection_names.append(f'file-{item_id}')
elif item_type == 'note':
# Note - always return full content as context
note = Notes.get_note_by_id(item_id)
note = await Notes.get_note_by_id(item_id)
if note and (
user_role == 'admin'
or note.user_id == user_id
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user_id,
resource_type='note',
resource_id=note.id,
@ -2099,11 +2099,11 @@ async def query_knowledge_files(
elif knowledge_ids:
# User specified specific KBs
for knowledge_id in knowledge_ids:
knowledge = Knowledges.get_knowledge_by_id(knowledge_id)
knowledge = await Knowledges.get_knowledge_by_id(knowledge_id)
if knowledge and (
user_role == 'admin'
or knowledge.user_id == user_id
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user_id,
resource_type='knowledge',
resource_id=knowledge.id,
@ -2114,7 +2114,7 @@ async def query_knowledge_files(
collection_names.append(knowledge_id)
else:
# No model knowledge and no specific IDs - search all accessible KBs
result = Knowledges.search_knowledge_bases(
result = await Knowledges.search_knowledge_bases(
user_id,
filter={
'query': '',
@ -2193,7 +2193,7 @@ async def query_knowledge_bases(
from open_webui.retrieval.vector.factory import VECTOR_DB_CLIENT
user_id = __user__.get('id')
user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id)]
user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)]
query_embedding = await __request__.app.state.EMBEDDING_FUNCTION(query)
# Min-heap of (distance, knowledge_base_id) - only holds top `count` results
@ -2203,7 +2203,7 @@ async def query_knowledge_bases(
page_size = 100
while True:
accessible_knowledge_bases = Knowledges.search_knowledge_bases(
accessible_knowledge_bases = await Knowledges.search_knowledge_bases(
user_id,
filter={'user_id': user_id, 'group_ids': user_group_ids},
skip=page_offset,
@ -2247,7 +2247,7 @@ async def query_knowledge_bases(
matching_knowledge_bases = []
for distance, knowledge_base_id in sorted_results:
knowledge_base = Knowledges.get_knowledge_by_id(knowledge_base_id)
knowledge_base = await Knowledges.get_knowledge_by_id(knowledge_base_id)
if knowledge_base:
matching_knowledge_bases.append(
{
@ -2295,7 +2295,7 @@ async def view_skill(
user_id = __user__.get('id')
# Direct DB lookup by id (case-insensitive since IDs are stored lowercase)
skill = Skills.get_skill_by_id(id.lower())
skill = await Skills.get_skill_by_id(id.lower())
if not skill or not skill.is_active:
return json.dumps({'error': f"Skill '{id}' not found"})
@ -2303,8 +2303,8 @@ async def view_skill(
# Check user access
user_role = __user__.get('role', 'user')
if user_role != 'admin' and skill.user_id != user_id:
user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id)]
if not AccessGrants.has_access(
user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)]
if not await AccessGrants.has_access(
user_id=user_id,
resource_type='skill',
resource_id=skill.id,
@ -2393,7 +2393,7 @@ async def tasks(
if tasks is None:
# Read-only - return current list
all_tasks = Chats.get_chat_tasks_by_id(__chat_id__)
all_tasks = await Chats.get_chat_tasks_by_id(__chat_id__)
elif overwrite:
# Full replacement - validate and write
all_tasks = []
@ -2417,7 +2417,7 @@ async def tasks(
)
else:
# Partial update - merge by id
existing_tasks = Chats.get_chat_tasks_by_id(__chat_id__)
existing_tasks = await Chats.get_chat_tasks_by_id(__chat_id__)
existing_by_id = {t['id']: t for t in existing_tasks}
seen_ids = set()
@ -2460,7 +2460,7 @@ async def tasks(
# Persist to DB and emit (skip for read-only)
if tasks is not None:
Chats.update_chat_tasks_by_id(__chat_id__, all_tasks)
await Chats.update_chat_tasks_by_id(__chat_id__, all_tasks)
if __event_emitter__:
await __event_emitter__(
@ -2542,7 +2542,7 @@ async def create_automation(
from open_webui.utils.automations import validate_rrule, next_run_ns, next_n_runs_ns
user_id = __user__.get('id')
user = Users.get_user_by_id(user_id)
user = await Users.get_user_by_id(user_id)
if not user:
return json.dumps({'error': 'User not found'})
@ -2569,7 +2569,7 @@ async def create_automation(
is_active=True,
)
automation = Automations.insert(user_id, form, next_run_ns(rrule, tz=tz))
automation = await Automations.insert(user_id, form, next_run_ns(rrule, tz=tz))
return json.dumps(
{
@ -2618,9 +2618,9 @@ async def update_automation(
from open_webui.utils.automations import validate_rrule, next_run_ns, next_n_runs_ns
user_id = __user__.get('id')
user = Users.get_user_by_id(user_id)
user = await Users.get_user_by_id(user_id)
automation = Automations.get_by_id(automation_id)
automation = await Automations.get_by_id(automation_id)
if not automation:
return json.dumps({'error': 'Automation not found'})
if automation.user_id != user_id:
@ -2650,7 +2650,7 @@ async def update_automation(
is_active=automation.is_active,
)
updated = Automations.update(automation_id, form, next_run_ns(new_rrule, tz=tz))
updated = await Automations.update(automation_id, form, next_run_ns(new_rrule, tz=tz))
return json.dumps(
{
@ -2693,9 +2693,9 @@ async def list_automations(
from open_webui.utils.automations import next_n_runs_ns
user_id = __user__.get('id')
user = Users.get_user_by_id(user_id)
user = await Users.get_user_by_id(user_id)
result = Automations.search_automations(
result = await Automations.search_automations(
user_id=user_id,
status=status,
skip=0,
@ -2753,16 +2753,16 @@ async def toggle_automation(
from open_webui.utils.automations import next_run_ns
user_id = __user__.get('id')
user = Users.get_user_by_id(user_id)
user = await Users.get_user_by_id(user_id)
automation = Automations.get_by_id(automation_id)
automation = await Automations.get_by_id(automation_id)
if not automation:
return json.dumps({'error': 'Automation not found'})
if automation.user_id != user_id:
return json.dumps({'error': 'Access denied'})
rrule = automation.data.get('rrule', '')
toggled = Automations.toggle(
toggled = await Automations.toggle(
automation_id,
next_run_ns(rrule, tz=user.timezone if user else None),
)
@ -2803,15 +2803,15 @@ async def delete_automation(
user_id = __user__.get('id')
automation = Automations.get_by_id(automation_id)
automation = await Automations.get_by_id(automation_id)
if not automation:
return json.dumps({'error': 'Automation not found'})
if automation.user_id != user_id:
return json.dumps({'error': 'Access denied'})
name = automation.name
AutomationRuns.delete_by_automation(automation_id)
Automations.delete(automation_id)
await AutomationRuns.delete_by_automation(automation_id)
await Automations.delete(automation_id)
return json.dumps(
{

View file

@ -11,7 +11,7 @@ from open_webui.models.access_grants import (
)
from open_webui.config import DEFAULT_USER_PERMISSIONS
from sqlalchemy.orm import Session
from sqlalchemy.ext.asyncio import AsyncSession
def fill_missing_permissions(permissions: dict[str, Any], default_permissions: dict[str, Any]) -> dict[str, Any]:
@ -28,10 +28,10 @@ def fill_missing_permissions(permissions: dict[str, Any], default_permissions: d
return permissions
def get_permissions(
async def get_permissions(
user_id: str,
default_permissions: dict[str, Any],
db: Session | None = None,
db: AsyncSession | None = None,
) -> dict[str, Any]:
"""
Get all permissions for a user by combining the permissions of all groups the user is a member of.
@ -53,7 +53,7 @@ def get_permissions(
permissions[key] = permissions[key] or value # Use the most permissive value (True > False)
return permissions
user_groups = Groups.get_groups_by_member_id(user_id, db=db)
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
# Deep copy default permissions to avoid modifying the original dict
permissions = json.loads(json.dumps(default_permissions))
@ -68,11 +68,11 @@ def get_permissions(
return permissions
def has_permission(
async def has_permission(
user_id: str,
permission_key: str,
default_permissions: dict[str, Any] = {},
db: Session | None = None,
db: AsyncSession | None = None,
) -> bool:
"""
Check if a user has a specific permission by checking the group permissions
@ -93,7 +93,7 @@ def has_permission(
permission_hierarchy = permission_key.split('.')
# Retrieve user group permissions
user_groups = Groups.get_groups_by_member_id(user_id, db=db)
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
for group in user_groups:
if get_permission(group.permissions or {}, permission_hierarchy):
@ -104,12 +104,12 @@ def has_permission(
return get_permission(default_permissions, permission_hierarchy)
def has_access(
async def has_access(
user_id: str,
permission: str = 'read',
access_grants: list | None = None,
user_group_ids: set[str] | None = None,
db: Session | None = None,
db: AsyncSession | None = None,
) -> bool:
"""
Check if a user has the specified permission using an in-memory access_grants list.
@ -126,7 +126,7 @@ def has_access(
return False
if user_group_ids is None:
user_groups = Groups.get_groups_by_member_id(user_id, db=db)
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
user_group_ids = {group.id for group in user_groups}
for grant in access_grants:
@ -144,7 +144,7 @@ def has_access(
return False
def has_connection_access(
async def has_connection_access(
user: UserModel,
connection: dict,
user_group_ids: set[str] | None = None,
@ -163,10 +163,10 @@ def has_connection_access(
return True
if user_group_ids is None:
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)}
user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)}
access_grants = (connection.get('config') or {}).get('access_grants', [])
return has_access(user.id, 'read', access_grants, user_group_ids)
return await has_access(user.id, 'read', access_grants, user_group_ids)
def migrate_access_control(data: dict, ac_key: str = 'access_control', grants_key: str = 'access_grants') -> None:
@ -210,13 +210,13 @@ def migrate_access_control(data: dict, ac_key: str = 'access_control', grants_ke
data.pop(ac_key, None)
def filter_allowed_access_grants(
async def filter_allowed_access_grants(
default_permissions: dict[str, Any],
user_id: str,
user_role: str,
access_grants: list,
public_permission_key: str,
db: Session | None = None,
db: AsyncSession | None = None,
) -> list:
"""
Checks if the user has the required permissions to grant access to a resource.
@ -228,7 +228,7 @@ def filter_allowed_access_grants(
# Check if user can share publicly
if (
has_public_read_access_grant(access_grants) or has_public_write_access_grant(access_grants)
) and not has_permission(
) and not await has_permission(
user_id,
public_permission_key,
default_permissions,
@ -246,7 +246,7 @@ def filter_allowed_access_grants(
]
# Strip individual user sharing if user lacks permission
if has_user_access_grant(access_grants) and not has_permission(
if has_user_access_grant(access_grants) and not await has_permission(
user_id,
'access_grants.allow_users',
default_permissions,
@ -257,7 +257,7 @@ def filter_allowed_access_grants(
return access_grants
def check_model_access(
async def check_model_access(
user: UserModel,
model_info,
bypass_filter: bool = False,
@ -270,7 +270,7 @@ def check_model_access(
Args:
user: The authenticated user.
model_info: The model record from Models.get_model_by_id(),
model_info: The model record from await Models.get_model_by_id(),
or None if the model is not registered.
bypass_filter: If True, skip all access checks (used by
internal callers and BYPASS_MODEL_ACCESS_CONTROL).
@ -284,10 +284,10 @@ def check_model_access(
if user.role == 'user':
from open_webui.models.access_grants import AccessGrants
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)}
user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)}
if not (
user.id == model_info.user_id
or AccessGrants.has_access(
or await AccessGrants.has_access(
user_id=user.id,
resource_type='model',
resource_id=model_info.id,

View file

@ -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:

View file

@ -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,

View file

@ -324,7 +324,7 @@ async def get_current_user(
# auth by api key
if token.startswith('sk-'):
user = get_current_user_by_api_key(request, token)
user = await get_current_user_by_api_key(request, token)
# Add user info to current span
if ENABLE_OTEL:
@ -356,7 +356,7 @@ async def get_current_user(
detail='Invalid token',
)
user = Users.get_user_by_id(data['id'])
user = await Users.get_user_by_id(data['id'])
if user is None:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
@ -382,10 +382,10 @@ async def get_current_user(
current_span.set_attribute('client.user.role', user.role)
current_span.set_attribute('client.auth.type', 'jwt')
# Refresh the user's last active timestamp asynchronously
# to prevent blocking the request
if background_tasks:
background_tasks.add_task(Users.update_last_active_by_id, user.id)
# Refresh the user's last active timestamp
# Fire-and-forget via asyncio.create_task to avoid blocking
import asyncio
asyncio.create_task(Users.update_last_active_by_id(user.id))
return user
else:
raise HTTPException(
@ -407,9 +407,9 @@ async def get_current_user(
raise e
def get_current_user_by_api_key(request, api_key: str):
async def get_current_user_by_api_key(request, api_key: str):
# Each function call manages its own short-lived session internally
user = Users.get_user_by_api_key(api_key)
user = await Users.get_user_by_api_key(api_key)
if user is None:
raise HTTPException(
@ -419,7 +419,7 @@ def get_current_user_by_api_key(request, api_key: str):
if not request.state.enable_api_keys or (
user.role != 'admin'
and not has_permission(
and not await has_permission(
user.id,
'features.api_keys',
request.app.state.config.USER_PERMISSIONS,
@ -438,7 +438,7 @@ def get_current_user_by_api_key(request, api_key: str):
current_span.set_attribute('client.user.role', user.role)
current_span.set_attribute('client.auth.type', 'api_key')
Users.update_last_active_by_id(user.id)
await Users.update_last_active_by_id(user.id)
return user
@ -460,7 +460,7 @@ def get_admin_user(user=Depends(get_current_user)):
return user
def create_admin_user(email: str, password: str, name: str = 'Admin'):
async def create_admin_user(email: str, password: str, name: str = 'Admin'):
"""
Create an admin user from environment variables.
Used for headless/automated deployments.
@ -470,14 +470,14 @@ def create_admin_user(email: str, password: str, name: str = 'Admin'):
if not email or not password:
return None
if Users.has_users():
if await Users.has_users():
log.debug('Users already exist, skipping admin creation')
return None
log.info(f'Creating admin account from environment variables: {email}')
try:
hashed = get_password_hash(password)
user = Auths.insert_new_auth(
user = await Auths.insert_new_auth(
email=email.lower(),
password=hashed,
name=name,

View file

@ -126,7 +126,7 @@ async def automation_worker_loop(app) -> None:
while True:
try:
with get_db() as db:
batch = Automations.claim_due(int(time.time_ns()), limit=10, db=db)
batch = await Automations.claim_due(int(time.time_ns()), limit=10, db=db)
if batch:
log.info(f'Claimed {len(batch)} due automation(s)')
for automation in batch:
@ -283,9 +283,9 @@ async def execute_automation(app, automation: AutomationModel) -> None:
(filters, model params, knowledge/RAG, tools, DB saves, webhooks).
"""
try:
user = Users.get_user_by_id(automation.user_id)
user = await Users.get_user_by_id(automation.user_id)
if not user:
_record_run(automation.id, 'error', error='User not found')
await _record_run(automation.id, 'error', error='User not found')
return
prompt = prompt_template(automation.data['prompt'], user)
@ -297,7 +297,7 @@ async def execute_automation(app, automation: AutomationModel) -> None:
assistant_msg_id = str(uuid4())
# Create the chat with user message (same structure as frontend)
chat = Chats.insert_new_chat(
chat = await Chats.insert_new_chat(
automation.user_id,
ChatForm(
chat={
@ -336,7 +336,7 @@ async def execute_automation(app, automation: AutomationModel) -> None:
)
if not chat:
_record_run(automation.id, 'error', error='Failed to create chat')
await _record_run(automation.id, 'error', error='Failed to create chat')
return
# Notify frontend to refresh chat list
@ -404,11 +404,11 @@ async def execute_automation(app, automation: AutomationModel) -> None:
room=f'user:{automation.user_id}',
)
_record_run(automation.id, 'success', chat_id=chat.id)
await _record_run(automation.id, 'success', chat_id=chat.id)
except Exception as e:
log.exception(f'Automation {automation.id} failed')
_record_run(automation.id, 'error', error=str(e)[:4000])
await _record_run(automation.id, 'error', error=str(e)[:4000])
####################
@ -416,7 +416,7 @@ async def execute_automation(app, automation: AutomationModel) -> None:
####################
def _record_run(
async def _record_run(
automation_id: str,
status: str,
chat_id: str = None,
@ -424,4 +424,4 @@ def _record_run(
):
"""Insert a run record into automation_run."""
with get_db() as db:
AutomationRuns.insert(automation_id, status, chat_id=chat_id, error=error, db=db)
await AutomationRuns.insert(automation_id, status, chat_id=chat_id, error=error, db=db)

View file

@ -72,7 +72,7 @@ async def generate_direct_chat_completion(
session_id = metadata.get('session_id')
request_id = str(uuid.uuid4()) # Generate a unique request ID
event_caller = get_event_call(metadata)
event_caller = await get_event_call(metadata)
channel = f'{user_id}:{session_id}:{request_id}'
logging.info(f'WebSocket channel: {channel}')
@ -199,7 +199,7 @@ async def generate_chat_completion(
# Check if user has access to the model
if not bypass_filter and user.role == 'user':
try:
check_model_access(user, model)
await check_model_access(user, model)
except Exception as e:
raise e
@ -343,8 +343,8 @@ async def chat_completed(request: Request, form_data: dict, user: Any):
}
extra_params = {
'__event_emitter__': get_event_emitter(metadata),
'__event_call__': get_event_call(metadata),
'__event_emitter__': await get_event_emitter(metadata),
'__event_call__': await get_event_call(metadata),
'__user__': user.model_dump() if isinstance(user, UserModel) else {},
'__metadata__': metadata,
'__request__': request,
@ -352,8 +352,8 @@ async def chat_completed(request: Request, form_data: dict, user: Any):
}
try:
filter_ids = get_sorted_filter_ids(request, model, metadata.get('filter_ids', []))
filter_functions = Functions.get_functions_by_ids(filter_ids)
filter_ids = await get_sorted_filter_ids(request, model, metadata.get('filter_ids', []))
filter_functions = await Functions.get_functions_by_ids(filter_ids)
result, _ = await process_filter_functions(
request=request,

View file

@ -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':

View file

@ -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

View file

@ -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}')

View file

@ -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}')

View file

@ -44,7 +44,7 @@ def create_httpx_client(headers=None, timeout=None, auth=None):
return _build_httpx_client(headers=headers, timeout=timeout, auth=auth, verify=True)
def create_insecure_httpx_client(headers=None, timeout=None, auth=None):
async def create_insecure_httpx_client(headers=None, timeout=None, auth=None):
return _build_httpx_client(headers=headers, timeout=timeout, auth=auth, verify=False)

View file

@ -936,7 +936,7 @@ def apply_source_context_to_messages(
)
def process_tool_result(
async def process_tool_result(
request,
tool_function_name,
tool_result,
@ -1075,7 +1075,7 @@ def process_tool_result(
pass
tool_response.append(text)
elif item.get('type') in ['image', 'audio']:
file_url = get_file_url_from_base64(
file_url = await get_file_url_from_base64(
request,
f'data:{item.get("mimeType")};base64,{item.get("data", item.get("blob", ""))}',
{
@ -1304,7 +1304,7 @@ async def chat_completion_tools_handler(
except Exception as e:
tool_result = str(e)
tool_result, tool_result_files, tool_result_embeds = process_tool_result(
tool_result, tool_result_files, tool_result_embeds = await process_tool_result(
request,
tool_function_name,
tool_result,
@ -1602,7 +1602,7 @@ def get_images_from_messages(message_list):
return images
def get_image_urls(delta_images, request, metadata, user) -> list[str]:
async def get_image_urls(delta_images, request, metadata, user) -> list[str]:
if not isinstance(delta_images, list):
return []
@ -1616,21 +1616,21 @@ def get_image_urls(delta_images, request, metadata, user) -> list[str]:
continue
if url.startswith('data:image/png;base64'):
url = get_image_url_from_base64(request, url, metadata, user)
url = await get_image_url_from_base64(request, url, metadata, user)
image_urls.append(url)
return image_urls
def add_file_context(messages: list, chat_id: str, user) -> list:
async def add_file_context(messages: list, chat_id: str, user) -> list:
"""
Add file URLs to messages for native function calling.
"""
if not chat_id or chat_id.startswith('local:'):
return messages
chat = Chats.get_chat_by_id_and_user_id(chat_id, user.id)
chat = await Chats.get_chat_by_id_and_user_id(chat_id, user.id)
if not chat:
return messages
@ -1686,7 +1686,7 @@ async def chat_image_generation_handler(request: Request, form_data: dict, extra
if chat_id.startswith('local:'):
message_list = form_data.get('messages', [])
else:
chat = Chats.get_chat_by_id_and_user_id(chat_id, user.id)
chat = await Chats.get_chat_by_id_and_user_id(chat_id, user.id)
await __event_emitter__(
{
'type': 'status',
@ -2066,12 +2066,12 @@ async def convert_url_images_to_base64(form_data):
return form_data
def load_messages_from_db(chat_id: str, message_id: str) -> Optional[list[dict]]:
async def load_messages_from_db(chat_id: str, message_id: str) -> Optional[list[dict]]:
"""
Load the message chain from DB up to message_id,
keeping only LLM-relevant fields (role, content, output).
"""
messages_map = Chats.get_messages_map_by_chat_id(chat_id)
messages_map = await Chats.get_messages_map_by_chat_id(chat_id)
if not messages_map:
return None
@ -2149,7 +2149,7 @@ async def process_chat_payload(request, form_data, user, metadata, model):
parent_message_id = metadata.get('parent_message_id')
if chat_id and parent_message_id and not chat_id.startswith('local:'):
db_messages = load_messages_from_db(chat_id, parent_message_id)
db_messages = await load_messages_from_db(chat_id, parent_message_id)
if db_messages:
system_message = get_system_message(form_data.get('messages', []))
form_data['messages'] = [system_message, *db_messages] if system_message else db_messages
@ -2192,8 +2192,8 @@ async def process_chat_payload(request, form_data, user, metadata, model):
form_data = await convert_url_images_to_base64(form_data)
event_emitter = get_event_emitter(metadata)
event_caller = get_event_call(metadata)
event_emitter = await get_event_emitter(metadata)
event_caller = await get_event_call(metadata)
extra_params = {
'__event_emitter__': event_emitter,
@ -2231,14 +2231,14 @@ async def process_chat_payload(request, form_data, user, metadata, model):
chat_id = metadata.get('chat_id', None)
folder_id = None
if chat_id and user:
folder_id = Chats.get_chat_folder_id(chat_id, user.id)
folder_id = await Chats.get_chat_folder_id(chat_id, user.id)
# Fallback: use folder_id from metadata (temporary chats have no DB record)
if not folder_id:
folder_id = metadata.get('folder_id', None)
if folder_id and user:
folder = Folders.get_folder_by_id_and_user_id(folder_id, user.id)
folder = await Folders.get_folder_by_id_and_user_id(folder_id, user.id)
if folder and folder.data:
if 'system_prompt' in folder.data:
@ -2305,8 +2305,8 @@ async def process_chat_payload(request, form_data, user, metadata, model):
raise e
try:
filter_ids = get_sorted_filter_ids(request, model, metadata.get('filter_ids', []))
filter_functions = Functions.get_functions_by_ids(filter_ids)
filter_ids = await get_sorted_filter_ids(request, model, metadata.get('filter_ids', []))
filter_functions = await Functions.get_functions_by_ids(filter_ids)
form_data, flags = await process_filter_functions(
request=request,
@ -2399,12 +2399,13 @@ async def process_chat_payload(request, form_data, user, metadata, model):
if all_skill_ids:
from open_webui.models.skills import Skills as SkillsModel
accessible_skill_ids = {s.id for s in SkillsModel.get_skills_by_user_id(user.id, 'read')}
available_skills = [
s
for sid in all_skill_ids
if sid in accessible_skill_ids and (s := SkillsModel.get_skill_by_id(sid)) and s.is_active
]
accessible_skill_ids = {s.id for s in await SkillsModel.get_skills_by_user_id(user.id, 'read')}
available_skills = []
for sid in all_skill_ids:
if sid in accessible_skill_ids:
s = await SkillsModel.get_skill_by_id(sid)
if s and s.is_active:
available_skills.append(s)
skill_descriptions = ''
for skill in available_skills:
@ -2441,7 +2442,7 @@ async def process_chat_payload(request, form_data, user, metadata, model):
# Get folder files
folder_id = file_item.get('id', None)
if folder_id:
folder = Folders.get_folder_by_id_and_user_id(folder_id, user.id)
folder = await Folders.get_folder_by_id_and_user_id(folder_id, user.id)
if folder and folder.data and 'files' in folder.data:
files = [f for f in files if f.get('id', None) != folder_id]
files = [*files, *folder.data['files']]
@ -2495,7 +2496,7 @@ async def process_chat_payload(request, form_data, user, metadata, model):
continue
# Check access control for MCP server
if not has_connection_access(user, mcp_server_connection):
if not await has_connection_access(user, mcp_server_connection):
log.warning(f'Access denied to MCP server {server_id} for user {user.id}')
continue
@ -2556,7 +2557,7 @@ async def process_chat_payload(request, form_data, user, metadata, model):
tool_specs = await mcp_clients[server_id].list_tool_specs()
for tool_spec in tool_specs:
def make_tool_function(client, function_name):
async def make_tool_function(client, function_name):
async def tool_function(**kwargs):
return await client.call_tool(
function_name,
@ -2570,7 +2571,7 @@ async def process_chat_payload(request, form_data, user, metadata, model):
# Skip this function
continue
tool_function = make_tool_function(mcp_clients[server_id], tool_spec['name'])
tool_function = await make_tool_function(mcp_clients[server_id], tool_spec['name'])
mcp_tools_dict[f'{server_id}_{tool_spec["name"]}'] = {
'spec': {
@ -2664,8 +2665,8 @@ async def process_chat_payload(request, form_data, user, metadata, model):
if metadata.get('params', {}).get('function_calling') == 'native' and builtin_tools_enabled:
# Add file context to user messages
chat_id = metadata.get('chat_id')
form_data['messages'] = add_file_context(form_data.get('messages', []), chat_id, user)
builtin_tools = get_builtin_tools(
form_data['messages'] = await add_file_context(form_data.get('messages', []), chat_id, user)
builtin_tools = await get_builtin_tools(
request,
{
**extra_params,
@ -2755,7 +2756,7 @@ async def process_chat_payload(request, form_data, user, metadata, model):
return form_data, metadata, events
def get_event_emitter_and_caller(metadata):
async def get_event_emitter_and_caller(metadata):
event_emitter = None
event_caller = None
@ -2763,18 +2764,18 @@ def get_event_emitter_and_caller(metadata):
# It broadcasts to user:{user_id} room AND persists to DB,
# so it works for backend-initiated calls (automations, API).
if metadata.get('chat_id') and metadata.get('message_id'):
event_emitter = get_event_emitter(metadata)
event_emitter = await get_event_emitter(metadata)
# event_caller needs session_id — it calls back to a specific
# websocket session (used by direct tools, pyodide code interpreter).
if metadata.get('session_id') and metadata.get('chat_id') and metadata.get('message_id'):
event_caller = get_event_call(metadata)
event_caller = await get_event_call(metadata)
return event_emitter, event_caller
def build_chat_response_context(request, form_data, user, model, metadata, tasks, events):
event_emitter, event_caller = get_event_emitter_and_caller(metadata)
async def build_chat_response_context(request, form_data, user, model, metadata, tasks, events):
event_emitter, event_caller = await get_event_emitter_and_caller(metadata)
return {
'request': request,
'form_data': form_data,
@ -2862,7 +2863,7 @@ async def background_tasks_handler(ctx):
messages = []
if 'chat_id' in metadata and not metadata['chat_id'].startswith('local:'):
messages_map = Chats.get_messages_map_by_chat_id(metadata['chat_id'])
messages_map = await Chats.get_messages_map_by_chat_id(metadata['chat_id'])
message = messages_map.get(metadata['message_id']) if messages_map else None
message_list = get_message_list(messages_map, metadata['message_id'])
@ -2942,7 +2943,7 @@ async def background_tasks_handler(ctx):
)
if not metadata.get('chat_id', '').startswith('local:'):
Chats.upsert_message_to_chat_by_id_and_message_id(
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
{
@ -2995,7 +2996,7 @@ async def background_tasks_handler(ctx):
if not title:
title = messages[0].get('content', user_message)
Chats.update_chat_title_by_id(metadata['chat_id'], title)
await Chats.update_chat_title_by_id(metadata['chat_id'], title)
await event_emitter(
{
@ -3007,7 +3008,7 @@ async def background_tasks_handler(ctx):
if title == None and len(messages) == 2 and (not messages_map or len(messages_map) <= 2):
title = messages[0].get('content', user_message)
Chats.update_chat_title_by_id(metadata['chat_id'], title)
await Chats.update_chat_title_by_id(metadata['chat_id'], title)
await event_emitter(
{
@ -3041,7 +3042,7 @@ async def background_tasks_handler(ctx):
try:
tags = json.loads(tags_string).get('tags', [])
Chats.update_chat_tags_by_id(metadata['chat_id'], tags, user)
await Chats.update_chat_tags_by_id(metadata['chat_id'], tags, user)
await event_emitter(
{
@ -3076,7 +3077,7 @@ async def non_streaming_chat_response_handler(response, ctx):
else:
error = str(error)
Chats.upsert_message_to_chat_by_id_and_message_id(
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
{
@ -3092,7 +3093,7 @@ async def non_streaming_chat_response_handler(response, ctx):
)
if 'selected_model_id' in response_data:
Chats.upsert_message_to_chat_by_id_and_message_id(
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
{
@ -3112,7 +3113,7 @@ async def non_streaming_chat_response_handler(response, ctx):
}
)
title = Chats.get_chat_title_by_id(metadata['chat_id'])
title = await Chats.get_chat_title_by_id(metadata['chat_id'])
# Use output from backend if provided (OR-compliant backends),
# otherwise generate from response content
@ -3143,7 +3144,7 @@ async def non_streaming_chat_response_handler(response, ctx):
# Save message in the database
usage = normalize_usage(response_data.get('usage', {}) or {})
Chats.upsert_message_to_chat_by_id_and_message_id(
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
{
@ -3156,8 +3157,8 @@ async def non_streaming_chat_response_handler(response, ctx):
)
# Send a webhook notification if the user is not active
if request.app.state.config.ENABLE_USER_WEBHOOKS and not Users.is_user_active(user.id):
webhook_url = Users.get_user_webhook_url_by_id(user.id)
if request.app.state.config.ENABLE_USER_WEBHOOKS and not await Users.is_user_active(user.id):
webhook_url = await Users.get_user_webhook_url_by_id(user.id)
if webhook_url:
await post_webhook(
request.app.state.WEBUI_NAME,
@ -3211,8 +3212,8 @@ async def streaming_chat_response_handler(response, ctx):
}
filter_functions = [
Functions.get_function_by_id(filter_id)
for filter_id in get_sorted_filter_ids(request, model, metadata.get('filter_ids', []))
await Functions.get_function_by_id(filter_id)
for filter_id in await get_sorted_filter_ids(request, model, metadata.get('filter_ids', []))
]
# Standard streaming response handler
@ -3447,7 +3448,7 @@ async def streaming_chat_response_handler(response, ctx):
return output, end_flag
message = Chats.get_message_by_id_and_message_id(metadata['chat_id'], metadata['message_id'])
message = await Chats.get_message_by_id_and_message_id(metadata['chat_id'], metadata['message_id'])
tool_calls = []
@ -3509,7 +3510,7 @@ async def streaming_chat_response_handler(response, ctx):
)
# Save message in the database
Chats.upsert_message_to_chat_by_id_and_message_id(
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
{
@ -3579,7 +3580,7 @@ async def streaming_chat_response_handler(response, ctx):
if 'selected_model_id' in data:
model_id = data['selected_model_id']
Chats.upsert_message_to_chat_by_id_and_message_id(
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
{
@ -3645,7 +3646,7 @@ async def streaming_chat_response_handler(response, ctx):
error = data.get('error', {})
if error:
try:
Chats.upsert_message_to_chat_by_id_and_message_id(
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
{
@ -3762,10 +3763,10 @@ async def streaming_chat_response_handler(response, ctx):
}
)
image_urls = get_image_urls(delta.get('images', []), request, metadata, user)
image_urls = await get_image_urls(delta.get('images', []), request, metadata, user)
if image_urls:
image_file_list = [{'type': 'image', 'url': url} for url in image_urls]
message_files = Chats.add_message_files_by_id_and_message_id(
message_files = await Chats.add_message_files_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
image_file_list,
@ -3847,7 +3848,7 @@ async def streaming_chat_response_handler(response, ctx):
)
if ENABLE_CHAT_RESPONSE_BASE64_IMAGE_URL_CONVERSION:
value = convert_markdown_base64_images(
value = await convert_markdown_base64_images(
request,
value,
{
@ -3963,7 +3964,7 @@ async def streaming_chat_response_handler(response, ctx):
if ENABLE_REALTIME_CHAT_SAVE:
# Save message in the database
Chats.upsert_message_to_chat_by_id_and_message_id(
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
{
@ -4184,7 +4185,7 @@ async def streaming_chat_response_handler(response, ctx):
)
else:
tool_function = get_updated_tool_function(
tool_function = await get_updated_tool_function(
function=tool['callable'],
extra_params={
'__messages__': form_data.get('messages', []),
@ -4197,7 +4198,7 @@ async def streaming_chat_response_handler(response, ctx):
except Exception as e:
tool_result = str(e)
tool_result, tool_result_files, tool_result_embeds = process_tool_result(
tool_result, tool_result_files, tool_result_embeds = await process_tool_result(
request,
tool_function_name,
tool_result,
@ -4487,7 +4488,7 @@ async def streaming_chat_response_handler(response, ctx):
BLOCKED_MODULES = {CODE_INTERPRETER_BLOCKED_MODULES}
_real_import = builtins.__import__
def restricted_import(name, globals=None, locals=None, fromlist=(), level=0):
async def restricted_import(name, globals=None, locals=None, fromlist=(), level=0):
if name.split('.')[0] in BLOCKED_MODULES:
importer_name = globals.get('__name__') if globals else None
if importer_name == '__main__':
@ -4541,7 +4542,7 @@ async def streaming_chat_response_handler(response, ctx):
stdoutLines = stdout.split('\n')
for idx, line in enumerate(stdoutLines):
if re.match(r'data:image/\w+;base64', line):
image_url = get_image_url_from_base64(
image_url = await get_image_url_from_base64(
request,
line,
metadata,
@ -4558,7 +4559,7 @@ async def streaming_chat_response_handler(response, ctx):
resultLines = result.split('\n')
for idx, line in enumerate(resultLines):
if re.match(r'data:image/\w+;base64', line):
image_url = get_image_url_from_base64(
image_url = await get_image_url_from_base64(
request,
line,
metadata,
@ -4623,7 +4624,7 @@ async def streaming_chat_response_handler(response, ctx):
if item.get('status') == 'in_progress':
item['status'] = 'completed'
title = Chats.get_chat_title_by_id(metadata['chat_id'])
title = await Chats.get_chat_title_by_id(metadata['chat_id'])
data = {
'done': True,
'content': serialize_output(output),
@ -4634,7 +4635,7 @@ async def streaming_chat_response_handler(response, ctx):
if not ENABLE_REALTIME_CHAT_SAVE:
# Save message in the database
Chats.upsert_message_to_chat_by_id_and_message_id(
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
{
@ -4645,21 +4646,21 @@ async def streaming_chat_response_handler(response, ctx):
},
)
elif usage:
Chats.upsert_message_to_chat_by_id_and_message_id(
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
{'done': True, 'usage': usage},
)
else:
Chats.upsert_message_to_chat_by_id_and_message_id(
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
{'done': True},
)
# Send a webhook notification if the user is not active
if request.app.state.config.ENABLE_USER_WEBHOOKS and not Users.is_user_active(user.id):
webhook_url = Users.get_user_webhook_url_by_id(user.id)
if request.app.state.config.ENABLE_USER_WEBHOOKS and not await Users.is_user_active(user.id):
webhook_url = await Users.get_user_webhook_url_by_id(user.id)
if webhook_url:
await post_webhook(
request.app.state.WEBUI_NAME,
@ -4687,7 +4688,7 @@ async def streaming_chat_response_handler(response, ctx):
if not ENABLE_REALTIME_CHAT_SAVE:
# Save message in the database
Chats.upsert_message_to_chat_by_id_and_message_id(
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
{
@ -4697,7 +4698,7 @@ async def streaming_chat_response_handler(response, ctx):
},
)
else:
Chats.upsert_message_to_chat_by_id_and_message_id(
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
{'done': True},

View file

@ -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,

View file

@ -700,7 +700,7 @@ class OAuthClientManager:
"""
try:
# Get the OAuth session
session = OAuthSessions.get_session_by_provider_and_user_id(client_id, user_id)
session = await OAuthSessions.get_session_by_provider_and_user_id(client_id, user_id)
if not session:
log.warning(f'No OAuth session found for user {user_id}, client_id {client_id}')
return None
@ -714,7 +714,7 @@ class OAuthClientManager:
log.warning(
f'Token refresh failed for user {user_id}, client_id {session.provider}, deleting session {session.id}'
)
OAuthSessions.delete_session_by_id(session.id)
await OAuthSessions.delete_session_by_id(session.id)
return None
return session.token
@ -738,7 +738,7 @@ class OAuthClientManager:
if refreshed_token:
# Update the session with new token data
session = OAuthSessions.update_session_by_id(session.id, refreshed_token)
session = await OAuthSessions.update_session_by_id(session.id, refreshed_token)
log.info(f'Successfully refreshed token for session {session.id}')
return session.token
else:
@ -884,12 +884,12 @@ class OAuthClientManager:
token['expires_at'] = datetime.now().timestamp() + token['expires_in']
# Clean up any existing sessions for this user/client_id first
sessions = OAuthSessions.get_sessions_by_user_id(user_id)
sessions = await OAuthSessions.get_sessions_by_user_id(user_id)
for session in sessions:
if session.provider == client_id:
OAuthSessions.delete_session_by_id(session.id)
await OAuthSessions.delete_session_by_id(session.id)
session = OAuthSessions.create_session(
session = await OAuthSessions.create_session(
user_id=user_id,
provider=client_id,
token=token,
@ -963,7 +963,7 @@ class OAuthManager:
"""
try:
# Get the OAuth session
session = OAuthSessions.get_session_by_id_and_user_id(session_id, user_id)
session = await OAuthSessions.get_session_by_id_and_user_id(session_id, user_id)
if not session:
log.warning(f'No OAuth session found for user {user_id}, session {session_id}')
return None
@ -977,7 +977,7 @@ class OAuthManager:
log.warning(
f'Token refresh failed for user {user_id}, provider {session.provider}, deleting session {session.id}'
)
OAuthSessions.delete_session_by_id(session.id)
await OAuthSessions.delete_session_by_id(session.id)
return None
return session.token
@ -1002,7 +1002,7 @@ class OAuthManager:
if refreshed_token:
# Update the session with new token data
session = OAuthSessions.update_session_by_id(session.id, refreshed_token)
session = await OAuthSessions.update_session_by_id(session.id, refreshed_token)
log.info(f'Successfully refreshed token for session {session.id}')
return session.token
else:
@ -1102,16 +1102,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')
@ -1185,7 +1188,7 @@ class OAuthManager:
return role
def update_user_groups(self, user, user_data, default_permissions, db=None):
async def update_user_groups(self, user, user_data, default_permissions, db=None):
log.debug('Running OAUTH Group management')
oauth_claim = auth_manager_config.OAUTH_GROUPS_CLAIM
@ -1214,8 +1217,8 @@ class OAuthManager:
else:
user_oauth_groups = []
user_current_groups: list[GroupModel] = Groups.get_groups_by_member_id(user.id, db=db)
all_available_groups: list[GroupModel] = Groups.get_all_groups(db=db)
user_current_groups: list[GroupModel] = await Groups.get_groups_by_member_id(user.id, db=db)
all_available_groups: list[GroupModel] = await Groups.get_all_groups(db=db)
# Create groups if they don't exist and creation is enabled
if auth_manager_config.ENABLE_OAUTH_GROUP_CREATION:
@ -1223,7 +1226,7 @@ class OAuthManager:
all_group_names = {g.name for g in all_available_groups}
groups_created = False
# Determine creator ID: Prefer admin, fallback to current user if no admin exists
admin_user = Users.get_super_admin_user()
admin_user = await Users.get_super_admin_user()
creator_id = admin_user.id if admin_user else user.id
log.debug(f'Using creator ID {creator_id} for potential group creation.')
@ -1238,7 +1241,7 @@ class OAuthManager:
data={'config': {'share': auth_manager_config.OAUTH_GROUP_DEFAULT_SHARE}},
)
# Use determined creator ID (admin or fallback to current user)
created_group = Groups.insert_new_group(creator_id, new_group_form, db=db)
created_group = await Groups.insert_new_group(creator_id, new_group_form, db=db)
if created_group:
log.info(
f"Successfully created group '{group_name}' with ID {created_group.id} using creator ID {creator_id}"
@ -1253,7 +1256,7 @@ class OAuthManager:
# Refresh the list of all available groups if any were created
if groups_created:
all_available_groups = Groups.get_all_groups(db=db)
all_available_groups = await Groups.get_all_groups(db=db)
log.debug('Refreshed list of all available groups after creation.')
log.debug(f'Oauth Groups claim: {oauth_claim}')
@ -1270,14 +1273,14 @@ class OAuthManager:
):
# Remove group from user
log.debug(f'Removing user from group {group_model.name} as it is no longer in their oauth groups')
Groups.remove_users_from_group(group_model.id, [user.id], db=db)
await Groups.remove_users_from_group(group_model.id, [user.id], db=db)
# In case a group is created, but perms are never assigned to the group by hitting "save"
group_permissions = group_model.permissions
if not group_permissions:
group_permissions = default_permissions
Groups.update_group_by_id(
await Groups.update_group_by_id(
id=group_model.id,
form_data=GroupUpdateForm(
name=group_model.name,
@ -1299,14 +1302,14 @@ class OAuthManager:
# Add user to group
log.debug(f'Adding user to group {group_model.name} as it was found in their oauth groups')
Groups.add_users_to_group(group_model.id, [user.id], db=db)
await Groups.add_users_to_group(group_model.id, [user.id], db=db)
# In case a group is created, but perms are never assigned to the group by hitting "save"
group_permissions = group_model.permissions
if not group_permissions:
group_permissions = default_permissions
Groups.update_group_by_id(
await Groups.update_group_by_id(
id=group_model.id,
form_data=GroupUpdateForm(
name=group_model.name,
@ -1487,20 +1490,20 @@ class OAuthManager:
raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED)
# Check if the user exists
user = Users.get_user_by_oauth_sub(provider, sub, db=db)
user = await Users.get_user_by_oauth_sub(provider, sub, db=db)
if not user:
# If the user does not exist, check if merging is enabled
if auth_manager_config.OAUTH_MERGE_ACCOUNTS_BY_EMAIL:
# Check if the user exists by email
user = Users.get_user_by_email(email, db=db)
user = await Users.get_user_by_email(email, db=db)
if user:
# Update the user with the new oauth sub
Users.update_user_oauth_by_id(user.id, provider, sub, db=db)
await Users.update_user_oauth_by_id(user.id, provider, sub, db=db)
if user:
determined_role = self.get_user_role(user, user_data)
determined_role = await self.get_user_role(user, user_data)
if user.role != determined_role:
Users.update_user_role_by_id(user.id, determined_role, db=db)
await Users.update_user_role_by_id(user.id, determined_role, db=db)
# Update the user object in memory as well,
# to avoid problems with the ENABLE_OAUTH_GROUP_MANAGEMENT check below
user.role = determined_role
@ -1510,7 +1513,7 @@ class OAuthManager:
if username_claim:
new_name = user_data.get(username_claim)
if new_name and new_name != user.name:
Users.update_user_by_id(user.id, {'name': new_name}, db=db)
await Users.update_user_by_id(user.id, {'name': new_name}, db=db)
user.name = new_name
log.debug(f'Updated name for user {user.email}')
@ -1519,13 +1522,13 @@ class OAuthManager:
if email_claim:
new_email = user_data.get(email_claim)
if new_email and new_email.lower() != user.email.lower():
existing_user = Users.get_user_by_email(new_email, db=db)
existing_user = await Users.get_user_by_email(new_email, db=db)
if existing_user:
log.error(
f'Cannot update email to {new_email} for user {user.id} because it is already taken.'
)
else:
Auths.update_email_by_id(user.id, new_email.lower(), db=db)
await Auths.update_email_by_id(user.id, new_email.lower(), db=db)
user.email = new_email.lower()
log.debug(f'Updated email for user {user.id}')
@ -1541,13 +1544,13 @@ class OAuthManager:
new_picture_url, token.get('access_token')
)
if processed_picture_url != user.profile_image_url:
Users.update_user_profile_image_url_by_id(user.id, processed_picture_url, db=db)
await Users.update_user_profile_image_url_by_id(user.id, processed_picture_url, db=db)
log.debug(f'Updated profile picture for user {user.email}')
else:
# If the user does not exist, check if signups are enabled
if auth_manager_config.ENABLE_OAUTH_SIGNUP:
# Check if an existing user with the same email already exists
existing_user = Users.get_user_by_email(email, db=db)
existing_user = await Users.get_user_by_email(email, db=db)
if existing_user:
raise HTTPException(400, detail=ERROR_MESSAGES.EMAIL_TAKEN)
@ -1567,16 +1570,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,
@ -1589,7 +1602,7 @@ class OAuthManager:
},
)
apply_default_group_assignment(request.app.state.config.DEFAULT_GROUP_ID, user.id, db=db)
await apply_default_group_assignment(request.app.state.config.DEFAULT_GROUP_ID, user.id, db=db)
else:
raise HTTPException(
@ -1602,7 +1615,7 @@ class OAuthManager:
expires_delta=parse_duration(auth_manager_config.JWT_EXPIRES_IN),
)
if auth_manager_config.ENABLE_OAUTH_GROUP_MANAGEMENT:
self.update_user_groups(
await self.update_user_groups(
user=user,
user_data=user_data,
default_permissions=request.app.state.config.USER_PERMISSIONS,
@ -1662,7 +1675,7 @@ class OAuthManager:
# Enforce max concurrent sessions per user/provider to prevent
# unbounded growth while allowing multi-device usage
sessions = OAuthSessions.get_sessions_by_user_id(user.id, db=db)
sessions = await OAuthSessions.get_sessions_by_user_id(user.id, db=db)
provider_sessions = sorted(
[session for session in sessions if session.provider == provider],
key=lambda session: session.created_at,
@ -1671,9 +1684,9 @@ class OAuthManager:
# Keep the newest sessions up to the limit, prune the rest
if len(provider_sessions) >= OAUTH_MAX_SESSIONS_PER_USER:
for old_session in provider_sessions[OAUTH_MAX_SESSIONS_PER_USER - 1 :]:
OAuthSessions.delete_session_by_id(old_session.id, db=db)
await OAuthSessions.delete_session_by_id(old_session.id, db=db)
session = OAuthSessions.create_session(
session = await OAuthSessions.create_session(
user_id=user.id,
provider=provider,
token=token,
@ -1834,7 +1847,7 @@ class OAuthManager:
# 8. Identify users to log out
users_to_logout = []
if sub:
user = Users.get_user_by_oauth_sub(matched_provider, sub, db=db)
user = await Users.get_user_by_oauth_sub(matched_provider, sub, db=db)
if user:
users_to_logout.append(user)
@ -1855,9 +1868,9 @@ class OAuthManager:
revoked_count = 0
for user in users_to_logout:
sessions = OAuthSessions.get_sessions_by_user_id(user.id, db=db)
sessions = await OAuthSessions.get_sessions_by_user_id(user.id, db=db)
for oauth_session in sessions:
OAuthSessions.delete_session_by_id(oauth_session.id, db=db)
await OAuthSessions.delete_session_by_id(oauth_session.id, db=db)
if redis:
revocation_key = f'{REDIS_KEY_PREFIX}:auth:user:{user.id}:revoked_at'

View file

@ -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:

View file

@ -38,7 +38,7 @@ class SentinelRedisProxy:
def _master(self):
return self._sentinel.master_for(self._service, **self._kw)
def __getattr__(self, item):
async def __getattr__(self, item):
master = self._master()
orig_attr = getattr(master, item)

View file

@ -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',

View file

@ -101,7 +101,7 @@ log = logging.getLogger(__name__)
# Let no function be called without need, and let what
# it yields justify the cost of running it.
def get_async_tool_function_and_apply_extra_params(function: Callable, extra_params: dict) -> Callable[..., Awaitable]:
async def get_async_tool_function_and_apply_extra_params(function: Callable, extra_params: dict) -> Callable[..., Awaitable]:
sig = inspect.signature(function)
extra_params = {k: v for k, v in extra_params.items() if k in sig.parameters}
partial_func = partial(function, **extra_params)
@ -138,13 +138,13 @@ def get_async_tool_function_and_apply_extra_params(function: Callable, extra_par
return new_function
def get_updated_tool_function(function: Callable, extra_params: dict):
async def get_updated_tool_function(function: Callable, extra_params: dict):
# Get the original function and merge updated params
__function__ = getattr(function, '__function__', None)
__extra_params__ = getattr(function, '__extra_params__', None)
if __function__ is not None and __extra_params__ is not None:
return get_async_tool_function_and_apply_extra_params(
return await get_async_tool_function_and_apply_extra_params(
__function__,
{**__extra_params__, **extra_params},
)
@ -160,16 +160,16 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr
tools_dict = {}
# Get user's group memberships for access control checks
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)}
user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)}
for tool_id in tool_ids:
tool = Tools.get_tool_by_id(tool_id)
tool = await Tools.get_tool_by_id(tool_id)
if tool:
# Check access control for local tools
if (
not (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL)
and tool.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='tool',
resource_id=tool.id,
@ -182,7 +182,7 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr
module = request.app.state.TOOLS.get(tool_id, None)
if module is None:
module, _ = load_tool_module_by_id(tool_id)
module, _ = await load_tool_module_by_id(tool_id)
request.app.state.TOOLS[tool_id] = module
__user__ = {
@ -191,11 +191,11 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr
# Set valves for the tool
if hasattr(module, 'valves') and hasattr(module, 'Valves'):
valves = Tools.get_tool_valves_by_id(tool_id) or {}
valves = await Tools.get_tool_valves_by_id(tool_id) or {}
module.valves = module.Valves(**valves)
if hasattr(module, 'UserValves'):
__user__['valves'] = module.UserValves( # type: ignore
**Tools.get_user_valves_by_id_and_user_id(tool_id, user.id)
**await Tools.get_user_valves_by_id_and_user_id(tool_id, user.id)
)
for spec in tool.specs:
@ -213,7 +213,7 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr
# convert to function that takes only model params and inserts custom params
function_name = spec['name']
tool_function = getattr(module, function_name)
callable = get_async_tool_function_and_apply_extra_params(
callable = await get_async_tool_function_and_apply_extra_params(
tool_function,
{
**extra_params,
@ -285,7 +285,7 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr
tool_server_connection = connections[tool_server_idx]
# Check access control for tool server
if not has_connection_access(user, tool_server_connection, user_group_ids):
if not await has_connection_access(user, tool_server_connection, user_group_ids):
log.warning(f'Access denied to tool server {server_id} for user {user.id}')
continue
@ -339,7 +339,7 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr
if metadata and metadata.get('message_id'):
headers[FORWARD_SESSION_INFO_HEADER_MESSAGE_ID] = metadata.get('message_id')
def make_tool_function(function_name, tool_server_data, headers):
async def make_tool_function(function_name, tool_server_data, headers):
async def tool_function(**kwargs):
return await execute_tool_server(
url=tool_server_data['url'],
@ -352,9 +352,9 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr
return tool_function
tool_function = make_tool_function(function_name, tool_server_data, headers)
tool_function = await make_tool_function(function_name, tool_server_data, headers)
callable = get_async_tool_function_and_apply_extra_params(
callable = await get_async_tool_function_and_apply_extra_params(
tool_function,
{},
)
@ -381,7 +381,7 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr
return tools_dict
def get_builtin_tools(
async def get_builtin_tools(
request: Request, extra_params: dict, features: dict = None, model: dict = None
) -> dict[str, dict]:
"""
@ -406,10 +406,10 @@ def get_builtin_tools(
# Helper to check user-level feature permission (admins always pass)
user = extra_params.get('__user__', {})
def has_user_permission(feature_key: str) -> bool:
async def has_user_permission(feature_key: str) -> bool:
if user.get('role') == 'admin':
return True
return has_permission(
return await has_permission(
user.get('id', ''),
f'features.{feature_key}',
request.app.state.config.USER_PERMISSIONS,
@ -461,7 +461,7 @@ def get_builtin_tools(
if (
is_builtin_tool_enabled('memory')
and (features.get('memory') or get_model_capability('memory', False))
and has_user_permission('memories')
and await has_user_permission('memories')
):
builtin_functions.extend(
[
@ -479,7 +479,7 @@ def get_builtin_tools(
and getattr(request.app.state.config, 'ENABLE_WEB_SEARCH', False)
and get_model_capability('web_search')
and features.get('web_search')
and has_user_permission('web_search')
and await has_user_permission('web_search')
):
builtin_functions.extend([search_web, fetch_url])
@ -489,7 +489,7 @@ def get_builtin_tools(
and getattr(request.app.state.config, 'ENABLE_IMAGE_GENERATION', False)
and get_model_capability('image_generation')
and features.get('image_generation')
and has_user_permission('image_generation')
and await has_user_permission('image_generation')
):
builtin_functions.append(generate_image)
if (
@ -497,7 +497,7 @@ def get_builtin_tools(
and getattr(request.app.state.config, 'ENABLE_IMAGE_EDIT', False)
and get_model_capability('image_generation')
and features.get('image_generation')
and has_user_permission('image_generation')
and await has_user_permission('image_generation')
):
builtin_functions.append(edit_image)
@ -507,7 +507,7 @@ def get_builtin_tools(
and getattr(request.app.state.config, 'ENABLE_CODE_INTERPRETER', True)
and get_model_capability('code_interpreter')
and features.get('code_interpreter')
and has_user_permission('code_interpreter')
and await has_user_permission('code_interpreter')
):
builtin_functions.append(execute_code)
@ -515,7 +515,7 @@ def get_builtin_tools(
if (
is_builtin_tool_enabled('notes')
and getattr(request.app.state.config, 'ENABLE_NOTES', False)
and has_user_permission('notes')
and await has_user_permission('notes')
):
builtin_functions.extend([search_notes, view_note, write_note, replace_note_content])
@ -523,7 +523,7 @@ def get_builtin_tools(
if (
is_builtin_tool_enabled('channels')
and getattr(request.app.state.config, 'ENABLE_CHANNELS', False)
and has_user_permission('channels')
and await has_user_permission('channels')
):
builtin_functions.extend(
[
@ -543,11 +543,11 @@ def get_builtin_tools(
builtin_functions.append(tasks)
# Automation tools - create and manage scheduled automations from chat
if is_builtin_tool_enabled('automations') and has_user_permission('automations'):
if is_builtin_tool_enabled('automations') and await has_user_permission('automations'):
builtin_functions.extend([create_automation, update_automation, list_automations, toggle_automation, delete_automation])
for func in builtin_functions:
callable = get_async_tool_function_and_apply_extra_params(
callable = await get_async_tool_function_and_apply_extra_params(
func,
{
'__request__': request,
@ -734,20 +734,31 @@ def get_tool_specs(tool_module: object) -> list[dict]:
return specs
def resolve_schema(schema, components):
def resolve_schema(schema, components, resolved_schemas=None):
"""
Recursively resolves a JSON schema using OpenAPI components.
"""
if not schema:
return {}
if resolved_schemas is None:
resolved_schemas = set()
if '$ref' in schema:
ref_path = schema['$ref']
schema_name = ref_path.split('/')[-1]
if schema_name in resolved_schemas:
# Avoid infinite recursion on circular references
return {}
resolved_schemas.add(schema_name)
ref_parts = ref_path.strip('#/').split('/')
resolved = components
for part in ref_parts[1:]: # Skip the initial 'components'
resolved = resolved.get(part, {})
return resolve_schema(resolved, components)
return resolve_schema(resolved, components, resolved_schemas)
resolved_schema = copy.deepcopy(schema)
@ -1013,8 +1024,8 @@ async def get_terminal_tools(
log.warning(f'Terminal server not found: {terminal_id}')
return {}
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)}
if not has_connection_access(user, connection, user_group_ids):
user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)}
if not await has_connection_access(user, connection, user_group_ids):
log.warning(f'Access denied to terminal {terminal_id} for user {user.id}')
return {}
@ -1066,7 +1077,7 @@ async def get_terminal_tools(
tool_spec.get('description', '') + f'\n\nThe current working directory is: {terminal_cwd}'
)
def make_tool_function(fn_name, srv_data, hdrs, cks):
async def make_tool_function(fn_name, srv_data, hdrs, cks):
async def tool_function(**kwargs):
return await execute_tool_server(
url=srv_data['url'],
@ -1079,8 +1090,8 @@ async def get_terminal_tools(
return tool_function
tool_function = make_tool_function(function_name, server_data, headers, cookies)
callable = get_async_tool_function_and_apply_extra_params(tool_function, {})
tool_function = await make_tool_function(function_name, server_data, headers, cookies)
callable = await get_async_tool_function_and_apply_extra_params(tool_function, {})
tools_dict[function_name] = {
'tool_id': f'terminal:{terminal_id}',

View file

@ -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

View file

@ -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

View file

@ -46,11 +46,11 @@
}}
>
<Tooltip
content={marked.parse(
content={DOMPurify.sanitize(marked.parse(
sanitizeResponseContent(
models[selectedModelIdx]?.info?.meta?.description ?? ''
).replaceAll('\n', '<br>')
)}
))}
placement="right"
>
<img
@ -97,13 +97,11 @@
<div
class="mt-0.5 text-base font-normal text-gray-500 dark:text-gray-400 line-clamp-3 markdown"
>
{@html DOMPurify.sanitize(
marked.parse(
sanitizeResponseContent(
models[selectedModelIdx]?.info?.meta?.description
).replaceAll('\n', '<br>')
)
)}
{@html DOMPurify.sanitize(marked.parse(
sanitizeResponseContent(
models[selectedModelIdx]?.info?.meta?.description
).replaceAll('\n', '<br>')
))}
</div>
{#if models[selectedModelIdx]?.info?.meta?.user}
<div class="mt-0.5 text-sm font-normal text-gray-400 dark:text-gray-500">

View file

@ -165,23 +165,21 @@
{#if models[selectedModelIdx]?.info?.meta?.description ?? null}
<Tooltip
className=" w-fit"
content={marked.parse(
content={DOMPurify.sanitize(marked.parse(
sanitizeResponseContent(
models[selectedModelIdx]?.info?.meta?.description ?? ''
).replaceAll('\n', '<br>')
)}
))}
placement="top"
>
<div
class="mt-0.5 px-2 text-sm font-normal text-gray-500 dark:text-gray-400 line-clamp-2 max-w-xl markdown"
>
{@html DOMPurify.sanitize(
marked.parse(
sanitizeResponseContent(
models[selectedModelIdx]?.info?.meta?.description ?? ''
).replaceAll('\n', '<br>')
)
)}
{@html DOMPurify.sanitize(marked.parse(
sanitizeResponseContent(
models[selectedModelIdx]?.info?.meta?.description ?? ''
).replaceAll('\n', '<br>')
))}
</div>
</Tooltip>

View file

@ -82,23 +82,27 @@
loading = true;
if (validateCommandString(command)) {
await onSubmit({
id: prompt?.id,
name,
command,
content,
tags: tags.map((tag) => tag.name),
access_grants: accessGrants,
commit_message: commitMessage || undefined,
is_production: isProduction
});
showEditModal = false;
commitMessage = '';
isProduction = true;
await loadHistory(true); // Reset and reload
// Select the newest version after saving
if (history.length > 0) {
selectedHistoryEntry = history[0];
try {
await onSubmit({
id: prompt?.id,
name,
command,
content,
tags: tags.map((tag) => tag.name),
access_grants: accessGrants,
commit_message: commitMessage || undefined,
is_production: isProduction
});
showEditModal = false;
commitMessage = '';
isProduction = true;
await loadHistory(true); // Reset and reload
// Select the newest version after saving
if (history.length > 0) {
selectedHistoryEntry = history[0];
}
} catch (error) {
toast.error(`${error}`);
}
} else {
toast.error(