mirror of
https://github.com/open-webui/open-webui.git
synced 2026-09-06 08:18:53 +00:00
Merge remote-tracking branch 'origin/dev' into fix/retrieval-collection-write-access
This commit is contained in:
commit
edf51e5c21
79 changed files with 4983 additions and 4549 deletions
|
|
@ -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)
|
||||
|
||||
####################################
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 = (
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}')
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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 = []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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}')
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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])
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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}')
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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',
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
{
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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':
|
||||
|
|
|
|||
|
|
@ -31,7 +31,7 @@ BASE64_IMAGE_URL_PREFIX = re.compile(r'data:image/\w+;base64,', re.IGNORECASE)
|
|||
MARKDOWN_IMAGE_URL_PATTERN = re.compile(r'!\[(.*?)\]\((.+?)\)', re.IGNORECASE)
|
||||
|
||||
|
||||
def get_image_base64_from_url(url: str) -> Optional[str]:
|
||||
async def get_image_base64_from_url(url: str) -> Optional[str]:
|
||||
try:
|
||||
if url.startswith('http'):
|
||||
# Validate URL to prevent SSRF attacks against local/private networks
|
||||
|
|
@ -44,7 +44,7 @@ def get_image_base64_from_url(url: str) -> Optional[str]:
|
|||
content_type = response.headers.get('Content-Type', 'image/png')
|
||||
return f'data:{content_type};base64,{encoded_string}'
|
||||
else:
|
||||
file = Files.get_file_by_id(url)
|
||||
file = await Files.get_file_by_id(url)
|
||||
|
||||
if not file:
|
||||
return None
|
||||
|
|
@ -64,13 +64,13 @@ def get_image_base64_from_url(url: str) -> Optional[str]:
|
|||
return None
|
||||
|
||||
|
||||
def get_image_url_from_base64(request, base64_image_string, metadata, user):
|
||||
async def get_image_url_from_base64(request, base64_image_string, metadata, user):
|
||||
if BASE64_IMAGE_URL_PREFIX.match(base64_image_string):
|
||||
image_url = ''
|
||||
# Extract base64 image data from the line
|
||||
image_data, content_type = get_image_data(base64_image_string)
|
||||
if image_data is not None:
|
||||
_, image_url = upload_image(
|
||||
_, image_url = await upload_image(
|
||||
request,
|
||||
image_data,
|
||||
content_type,
|
||||
|
|
@ -82,17 +82,26 @@ def get_image_url_from_base64(request, base64_image_string, metadata, user):
|
|||
return None
|
||||
|
||||
|
||||
def convert_markdown_base64_images(request, content: str, metadata, user):
|
||||
def replace(match):
|
||||
base64_string = match.group(2)
|
||||
MIN_REPLACEMENT_URL_LENGTH = 1024
|
||||
if len(base64_string) > MIN_REPLACEMENT_URL_LENGTH:
|
||||
url = get_image_url_from_base64(request, base64_string, metadata, user)
|
||||
if url:
|
||||
return f''
|
||||
return match.group(0)
|
||||
async def convert_markdown_base64_images(request, content: str, metadata, user):
|
||||
MIN_REPLACEMENT_URL_LENGTH = 1024
|
||||
result_parts = []
|
||||
last_end = 0
|
||||
|
||||
return MARKDOWN_IMAGE_URL_PATTERN.sub(replace, content)
|
||||
for match in MARKDOWN_IMAGE_URL_PATTERN.finditer(content):
|
||||
result_parts.append(content[last_end:match.start()])
|
||||
base64_string = match.group(2)
|
||||
if len(base64_string) > MIN_REPLACEMENT_URL_LENGTH:
|
||||
url = await get_image_url_from_base64(request, base64_string, metadata, user)
|
||||
if url:
|
||||
result_parts.append(f'')
|
||||
else:
|
||||
result_parts.append(match.group(0))
|
||||
else:
|
||||
result_parts.append(match.group(0))
|
||||
last_end = match.end()
|
||||
|
||||
result_parts.append(content[last_end:])
|
||||
return ''.join(result_parts)
|
||||
|
||||
|
||||
def load_b64_audio_data(b64_str):
|
||||
|
|
@ -110,7 +119,7 @@ def load_b64_audio_data(b64_str):
|
|||
return None, None
|
||||
|
||||
|
||||
def upload_audio(request, audio_data, content_type, metadata, user):
|
||||
async def upload_audio(request, audio_data, content_type, metadata, user):
|
||||
audio_format = mimetypes.guess_extension(content_type)
|
||||
file = UploadFile(
|
||||
file=io.BytesIO(audio_data),
|
||||
|
|
@ -119,7 +128,7 @@ def upload_audio(request, audio_data, content_type, metadata, user):
|
|||
'content-type': content_type,
|
||||
},
|
||||
)
|
||||
file_item = upload_file_handler(
|
||||
file_item = await upload_file_handler(
|
||||
request,
|
||||
file=file,
|
||||
metadata=metadata,
|
||||
|
|
@ -130,13 +139,13 @@ def upload_audio(request, audio_data, content_type, metadata, user):
|
|||
return url
|
||||
|
||||
|
||||
def get_audio_url_from_base64(request, base64_audio_string, metadata, user):
|
||||
async def get_audio_url_from_base64(request, base64_audio_string, metadata, user):
|
||||
if 'data:audio/wav;base64' in base64_audio_string:
|
||||
audio_url = ''
|
||||
# Extract base64 audio data from the line
|
||||
audio_data, content_type = load_b64_audio_data(base64_audio_string)
|
||||
if audio_data is not None:
|
||||
audio_url = upload_audio(
|
||||
audio_url = await upload_audio(
|
||||
request,
|
||||
audio_data,
|
||||
content_type,
|
||||
|
|
@ -147,16 +156,16 @@ def get_audio_url_from_base64(request, base64_audio_string, metadata, user):
|
|||
return None
|
||||
|
||||
|
||||
def get_file_url_from_base64(request, base64_file_string, metadata, user):
|
||||
async def get_file_url_from_base64(request, base64_file_string, metadata, user):
|
||||
if BASE64_IMAGE_URL_PREFIX.match(base64_file_string):
|
||||
return get_image_url_from_base64(request, base64_file_string, metadata, user)
|
||||
return await get_image_url_from_base64(request, base64_file_string, metadata, user)
|
||||
elif 'data:audio/wav;base64' in base64_file_string:
|
||||
return get_audio_url_from_base64(request, base64_file_string, metadata, user)
|
||||
return await get_audio_url_from_base64(request, base64_file_string, metadata, user)
|
||||
return None
|
||||
|
||||
|
||||
def get_image_base64_from_file_id(id: str) -> Optional[str]:
|
||||
file = Files.get_file_by_id(id)
|
||||
async def get_image_base64_from_file_id(id: str) -> Optional[str]:
|
||||
file = await Files.get_file_by_id(id)
|
||||
if not file:
|
||||
return None
|
||||
|
||||
|
|
|
|||
|
|
@ -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}')
|
||||
|
|
|
|||
|
|
@ -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}')
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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'
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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',
|
||||
|
|
|
|||
|
|
@ -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}',
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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">
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue